0.29.1原版
This commit is contained in:
@@ -0,0 +1,48 @@
|
||||
package v1
|
||||
|
||||
// PublicMethods defines API endpoints that don't require authentication.
|
||||
// All other endpoints require a valid session or access token.
|
||||
//
|
||||
// This is the SINGLE SOURCE OF TRUTH for public endpoints.
|
||||
// Both Connect interceptor and gRPC-Gateway interceptor use this map.
|
||||
//
|
||||
// Format: Full gRPC procedure path as returned by req.Spec().Procedure (Connect)
|
||||
// or info.FullMethod (gRPC interceptor).
|
||||
var PublicMethods = map[string]struct{}{
|
||||
// Auth Service - login/token endpoints must be accessible without auth
|
||||
"/memos.api.v1.AuthService/SignIn": {},
|
||||
"/memos.api.v1.AuthService/RefreshToken": {}, // Token refresh uses cookie, must be accessible when access token expired
|
||||
|
||||
// Instance Service - needed before login to show instance info
|
||||
"/memos.api.v1.InstanceService/GetInstanceProfile": {},
|
||||
"/memos.api.v1.InstanceService/GetInstanceSetting": {},
|
||||
"/memos.api.v1.InstanceService/BatchGetInstanceSettings": {},
|
||||
|
||||
// User Service - public user profiles and stats
|
||||
"/memos.api.v1.UserService/CreateUser": {}, // Allow first user registration
|
||||
"/memos.api.v1.UserService/GetUser": {},
|
||||
"/memos.api.v1.UserService/BatchGetUsers": {},
|
||||
"/memos.api.v1.UserService/GetUserAvatar": {},
|
||||
"/memos.api.v1.UserService/GetUserStats": {},
|
||||
"/memos.api.v1.UserService/ListAllUserStats": {},
|
||||
|
||||
// Identity Provider Service - SSO buttons on login page
|
||||
"/memos.api.v1.IdentityProviderService/ListIdentityProviders": {},
|
||||
|
||||
// Memo Service - public memos (visibility filtering done in service layer)
|
||||
"/memos.api.v1.MemoService/GetMemo": {},
|
||||
"/memos.api.v1.MemoService/ListMemos": {},
|
||||
"/memos.api.v1.MemoService/ListMemoComments": {},
|
||||
"/memos.api.v1.MemoService/GetLinkMetadata": {},
|
||||
"/memos.api.v1.MemoService/BatchGetLinkMetadata": {},
|
||||
|
||||
// Memo sharing - share-token endpoints require no authentication
|
||||
"/memos.api.v1.MemoService/GetMemoByShare": {},
|
||||
}
|
||||
|
||||
// IsPublicMethod checks if a procedure path is public (no authentication required).
|
||||
// Returns true for public methods, false for protected methods.
|
||||
func IsPublicMethod(procedure string) bool {
|
||||
_, ok := PublicMethods[procedure]
|
||||
return ok
|
||||
}
|
||||
@@ -0,0 +1,90 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
)
|
||||
|
||||
// TestPublicMethodsArePublic verifies that methods in PublicMethods are recognized as public.
|
||||
func TestPublicMethodsArePublic(t *testing.T) {
|
||||
publicMethods := []string{
|
||||
// Auth Service
|
||||
"/memos.api.v1.AuthService/SignIn",
|
||||
"/memos.api.v1.AuthService/RefreshToken",
|
||||
// Instance Service
|
||||
"/memos.api.v1.InstanceService/GetInstanceProfile",
|
||||
"/memos.api.v1.InstanceService/GetInstanceSetting",
|
||||
"/memos.api.v1.InstanceService/BatchGetInstanceSettings",
|
||||
// User Service
|
||||
"/memos.api.v1.UserService/CreateUser",
|
||||
"/memos.api.v1.UserService/GetUser",
|
||||
"/memos.api.v1.UserService/BatchGetUsers",
|
||||
"/memos.api.v1.UserService/GetUserAvatar",
|
||||
"/memos.api.v1.UserService/GetUserStats",
|
||||
"/memos.api.v1.UserService/ListAllUserStats",
|
||||
// Identity Provider Service
|
||||
"/memos.api.v1.IdentityProviderService/ListIdentityProviders",
|
||||
// Memo Service
|
||||
"/memos.api.v1.MemoService/GetMemo",
|
||||
"/memos.api.v1.MemoService/ListMemos",
|
||||
"/memos.api.v1.MemoService/GetLinkMetadata",
|
||||
"/memos.api.v1.MemoService/BatchGetLinkMetadata",
|
||||
}
|
||||
|
||||
for _, method := range publicMethods {
|
||||
t.Run(method, func(t *testing.T) {
|
||||
assert.True(t, IsPublicMethod(method), "Expected %s to be public", method)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestProtectedMethodsRequireAuth verifies that non-public methods are recognized as protected.
|
||||
func TestProtectedMethodsRequireAuth(t *testing.T) {
|
||||
protectedMethods := []string{
|
||||
// Auth Service - logout and get current user require auth
|
||||
"/memos.api.v1.AuthService/SignOut",
|
||||
"/memos.api.v1.AuthService/GetCurrentUser",
|
||||
// Instance Service - admin operations
|
||||
"/memos.api.v1.InstanceService/UpdateInstanceSetting",
|
||||
"/memos.api.v1.InstanceService/TestInstanceEmailSetting",
|
||||
// User Service - modification operations
|
||||
"/memos.api.v1.UserService/ListUsers",
|
||||
"/memos.api.v1.UserService/UpdateUser",
|
||||
"/memos.api.v1.UserService/DeleteUser",
|
||||
// Memo Service - write operations
|
||||
"/memos.api.v1.MemoService/CreateMemo",
|
||||
"/memos.api.v1.MemoService/UpdateMemo",
|
||||
"/memos.api.v1.MemoService/DeleteMemo",
|
||||
// Attachment Service - write operations
|
||||
"/memos.api.v1.AttachmentService/CreateAttachment",
|
||||
"/memos.api.v1.AttachmentService/DeleteAttachment",
|
||||
// Shortcut Service
|
||||
"/memos.api.v1.ShortcutService/CreateShortcut",
|
||||
"/memos.api.v1.ShortcutService/ListShortcuts",
|
||||
"/memos.api.v1.ShortcutService/UpdateShortcut",
|
||||
"/memos.api.v1.ShortcutService/DeleteShortcut",
|
||||
}
|
||||
|
||||
for _, method := range protectedMethods {
|
||||
t.Run(method, func(t *testing.T) {
|
||||
assert.False(t, IsPublicMethod(method), "Expected %s to require auth", method)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestUnknownMethodsRequireAuth verifies that unknown methods default to requiring auth.
|
||||
func TestUnknownMethodsRequireAuth(t *testing.T) {
|
||||
unknownMethods := []string{
|
||||
"/unknown.Service/Method",
|
||||
"/memos.api.v1.UnknownService/Method",
|
||||
"",
|
||||
"invalid",
|
||||
}
|
||||
|
||||
for _, method := range unknownMethods {
|
||||
t.Run(method, func(t *testing.T) {
|
||||
assert.False(t, IsPublicMethod(method), "Unknown method %q should require auth", method)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,240 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"mime"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
"github.com/usememos/memos/internal/ai"
|
||||
"github.com/usememos/memos/internal/ai/audiollm"
|
||||
audiollmgemini "github.com/usememos/memos/internal/ai/audiollm/gemini"
|
||||
"github.com/usememos/memos/internal/ai/stt"
|
||||
sttopenai "github.com/usememos/memos/internal/ai/stt/openai"
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
)
|
||||
|
||||
const (
|
||||
maxTranscriptionAudioSizeBytes = 25 * MebiByte
|
||||
maxTranscriptionFilenameLength = 255
|
||||
)
|
||||
|
||||
var supportedTranscriptionContentTypes = map[string]bool{
|
||||
"audio/aac": true,
|
||||
"audio/aiff": true,
|
||||
"audio/flac": true,
|
||||
"audio/mpeg": true,
|
||||
"audio/mp3": true,
|
||||
"audio/mp4": true,
|
||||
"audio/mpga": true,
|
||||
"audio/ogg": true,
|
||||
"audio/wav": true,
|
||||
"audio/x-wav": true,
|
||||
"audio/x-flac": true,
|
||||
"audio/x-m4a": true,
|
||||
"audio/webm": true,
|
||||
"video/mp4": true,
|
||||
"video/mpeg": true,
|
||||
"video/webm": true,
|
||||
}
|
||||
|
||||
// Transcribe transcribes an audio file using an instance AI provider.
|
||||
func (s *APIV1Service) Transcribe(ctx context.Context, request *v1pb.TranscribeRequest) (*v1pb.TranscribeResponse, error) {
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
|
||||
if request.Audio == nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "audio is required")
|
||||
}
|
||||
if request.Audio.GetUri() != "" {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "audio uri is not supported")
|
||||
}
|
||||
content := request.Audio.GetContent()
|
||||
if len(content) == 0 {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "audio content is required")
|
||||
}
|
||||
if len(content) > maxTranscriptionAudioSizeBytes {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "audio file is too large; maximum size is 25 MiB")
|
||||
}
|
||||
filename := strings.TrimSpace(request.Audio.GetFilename())
|
||||
if len(filename) > maxTranscriptionFilenameLength {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "filename is too long; maximum length is %d characters", maxTranscriptionFilenameLength)
|
||||
}
|
||||
contentType := strings.TrimSpace(request.Audio.GetContentType())
|
||||
if contentType == "" {
|
||||
contentType = http.DetectContentType(content)
|
||||
}
|
||||
if !isSupportedTranscriptionContentType(contentType) {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "audio content type %q is not supported", contentType)
|
||||
}
|
||||
|
||||
aiSetting, err := s.Store.GetInstanceAISetting(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get AI setting: %v", err)
|
||||
}
|
||||
persisted := aiSetting.GetTranscription()
|
||||
|
||||
providerID := persisted.GetProviderId()
|
||||
if providerID == "" {
|
||||
return nil, status.Errorf(codes.FailedPrecondition, "transcription is not configured")
|
||||
}
|
||||
|
||||
provider, err := s.resolveAIProvider(aiSetting, providerID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
model := persisted.GetModel()
|
||||
if model == "" {
|
||||
defaultModel, err := ai.DefaultTranscriptionModel(provider.Type)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "%v", err)
|
||||
}
|
||||
model = defaultModel
|
||||
}
|
||||
|
||||
var text string
|
||||
switch provider.Type {
|
||||
case ai.ProviderOpenAI:
|
||||
text, err = s.transcribeViaSTT(ctx, provider, persisted, model, content, filename, contentType)
|
||||
case ai.ProviderGemini:
|
||||
text, err = s.transcribeViaAudioLLM(ctx, provider, persisted, model, content, contentType)
|
||||
default:
|
||||
return nil, status.Errorf(codes.FailedPrecondition,
|
||||
"provider type %q is not supported for transcription", provider.Type)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to transcribe audio: %v", err)
|
||||
}
|
||||
return &v1pb.TranscribeResponse{Text: text}, nil
|
||||
}
|
||||
|
||||
func (*APIV1Service) transcribeViaSTT(
|
||||
ctx context.Context,
|
||||
provider ai.ProviderConfig,
|
||||
persisted *storepb.TranscriptionConfig,
|
||||
model string,
|
||||
content []byte,
|
||||
filename string,
|
||||
contentType string,
|
||||
) (string, error) {
|
||||
transcriber, err := sttopenai.New(provider, stt.ApplyOptions(nil))
|
||||
if err != nil {
|
||||
return "", errors.Wrap(err, "failed to create STT transcriber")
|
||||
}
|
||||
resp, err := transcriber.Transcribe(ctx, stt.Request{
|
||||
Audio: bytes.NewReader(content),
|
||||
Size: int64(len(content)),
|
||||
Filename: filename,
|
||||
ContentType: contentType,
|
||||
Model: model,
|
||||
Prompt: persisted.GetPrompt(),
|
||||
Language: persisted.GetLanguage(),
|
||||
})
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return resp.Text, nil
|
||||
}
|
||||
|
||||
func (*APIV1Service) transcribeViaAudioLLM(
|
||||
ctx context.Context,
|
||||
provider ai.ProviderConfig,
|
||||
persisted *storepb.TranscriptionConfig,
|
||||
model string,
|
||||
content []byte,
|
||||
contentType string,
|
||||
) (string, error) {
|
||||
m, err := audiollmgemini.New(provider, audiollm.ApplyOptions(nil))
|
||||
if err != nil {
|
||||
return "", errors.Wrap(err, "failed to create audio LLM")
|
||||
}
|
||||
resp, err := m.GenerateFromAudio(ctx, audiollm.Request{
|
||||
Audio: bytes.NewReader(content),
|
||||
Size: int64(len(content)),
|
||||
ContentType: contentType,
|
||||
Model: model,
|
||||
Instructions: buildTranscriptionInstructions(persisted.GetPrompt(), persisted.GetLanguage()),
|
||||
})
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if resp.FinishReason != audiollm.FinishStop {
|
||||
return "", errors.Errorf("transcription incomplete (finish reason: %s)", resp.FinishReason)
|
||||
}
|
||||
if strings.TrimSpace(resp.Text) == "" {
|
||||
return "", errors.New("transcription response did not include text")
|
||||
}
|
||||
return resp.Text, nil
|
||||
}
|
||||
|
||||
func buildTranscriptionInstructions(prompt, language string) string {
|
||||
parts := []string{
|
||||
"Transcribe the audio accurately. Return only the transcript text. " +
|
||||
"Do not summarize, explain, or add content that is not spoken.",
|
||||
}
|
||||
if language = strings.TrimSpace(language); language != "" {
|
||||
parts = append(parts, "The input language is "+language+".")
|
||||
}
|
||||
if prompt = strings.TrimSpace(prompt); prompt != "" {
|
||||
parts = append(parts, "Context and spelling hints:\n"+prompt)
|
||||
}
|
||||
return strings.Join(parts, "\n\n")
|
||||
}
|
||||
|
||||
func (*APIV1Service) resolveAIProvider(setting *storepb.InstanceAISetting, providerID string) (ai.ProviderConfig, error) {
|
||||
providers := make([]ai.ProviderConfig, 0, len(setting.GetProviders()))
|
||||
for _, provider := range setting.GetProviders() {
|
||||
if provider == nil {
|
||||
continue
|
||||
}
|
||||
providers = append(providers, convertAIProviderConfigFromStore(provider))
|
||||
}
|
||||
|
||||
provider, err := ai.FindProvider(providers, providerID)
|
||||
if err != nil {
|
||||
return ai.ProviderConfig{}, status.Errorf(codes.FailedPrecondition, "transcription provider is not configured")
|
||||
}
|
||||
return *provider, nil
|
||||
}
|
||||
|
||||
func convertAIProviderConfigFromStore(provider *storepb.AIProviderConfig) ai.ProviderConfig {
|
||||
return ai.ProviderConfig{
|
||||
ID: provider.GetId(),
|
||||
Title: provider.GetTitle(),
|
||||
Type: convertAIProviderTypeFromStore(provider.GetType()),
|
||||
Endpoint: provider.GetEndpoint(),
|
||||
APIKey: provider.GetApiKey(),
|
||||
}
|
||||
}
|
||||
|
||||
func convertAIProviderTypeFromStore(providerType storepb.AIProviderType) ai.ProviderType {
|
||||
switch providerType {
|
||||
case storepb.AIProviderType_OPENAI:
|
||||
return ai.ProviderOpenAI
|
||||
case storepb.AIProviderType_GEMINI:
|
||||
return ai.ProviderGemini
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
}
|
||||
|
||||
func isSupportedTranscriptionContentType(contentType string) bool {
|
||||
mediaType, _, err := mime.ParseMediaType(strings.TrimSpace(contentType))
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
mediaType = strings.ToLower(mediaType)
|
||||
return supportedTranscriptionContentTypes[mediaType]
|
||||
}
|
||||
@@ -0,0 +1,232 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/binary"
|
||||
"hash/crc32"
|
||||
"image"
|
||||
"image/color"
|
||||
"image/jpeg"
|
||||
"testing"
|
||||
|
||||
"github.com/disintegration/imaging"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestShouldStripExif(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
mimeType string
|
||||
expected bool
|
||||
}{
|
||||
{
|
||||
name: "JPEG should strip EXIF",
|
||||
mimeType: "image/jpeg",
|
||||
expected: true,
|
||||
},
|
||||
{
|
||||
name: "JPG should strip EXIF",
|
||||
mimeType: "image/jpg",
|
||||
expected: true,
|
||||
},
|
||||
{
|
||||
name: "TIFF should strip EXIF",
|
||||
mimeType: "image/tiff",
|
||||
expected: true,
|
||||
},
|
||||
{
|
||||
name: "WebP should strip EXIF",
|
||||
mimeType: "image/webp",
|
||||
expected: true,
|
||||
},
|
||||
{
|
||||
name: "HEIC should strip EXIF",
|
||||
mimeType: "image/heic",
|
||||
expected: true,
|
||||
},
|
||||
{
|
||||
name: "HEIF should strip EXIF",
|
||||
mimeType: "image/heif",
|
||||
expected: true,
|
||||
},
|
||||
{
|
||||
name: "PNG should not strip EXIF",
|
||||
mimeType: "image/png",
|
||||
expected: false,
|
||||
},
|
||||
{
|
||||
name: "GIF should not strip EXIF",
|
||||
mimeType: "image/gif",
|
||||
expected: false,
|
||||
},
|
||||
{
|
||||
name: "text file should not strip EXIF",
|
||||
mimeType: "text/plain",
|
||||
expected: false,
|
||||
},
|
||||
{
|
||||
name: "PDF should not strip EXIF",
|
||||
mimeType: "application/pdf",
|
||||
expected: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
result := shouldStripExif(tt.mimeType)
|
||||
assert.Equal(t, tt.expected, result)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestStripImageExif(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Create a simple test image
|
||||
img := image.NewRGBA(image.Rect(0, 0, 100, 100))
|
||||
// Fill with red color
|
||||
for y := 0; y < 100; y++ {
|
||||
for x := 0; x < 100; x++ {
|
||||
img.Set(x, y, color.RGBA{R: 255, G: 0, B: 0, A: 255})
|
||||
}
|
||||
}
|
||||
|
||||
// Encode as JPEG
|
||||
var buf bytes.Buffer
|
||||
err := jpeg.Encode(&buf, img, &jpeg.Options{Quality: 90})
|
||||
require.NoError(t, err)
|
||||
originalData := buf.Bytes()
|
||||
|
||||
t.Run("strip JPEG metadata", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
strippedData, err := stripImageExif(originalData, "image/jpeg")
|
||||
require.NoError(t, err)
|
||||
assert.NotEmpty(t, strippedData)
|
||||
|
||||
// Verify it's still a valid image
|
||||
decodedImg, err := imaging.Decode(bytes.NewReader(strippedData))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 100, decodedImg.Bounds().Dx())
|
||||
assert.Equal(t, 100, decodedImg.Bounds().Dy())
|
||||
})
|
||||
|
||||
t.Run("strip JPG metadata (alternate extension)", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
strippedData, err := stripImageExif(originalData, "image/jpg")
|
||||
require.NoError(t, err)
|
||||
assert.NotEmpty(t, strippedData)
|
||||
|
||||
// Verify it's still a valid image
|
||||
decodedImg, err := imaging.Decode(bytes.NewReader(strippedData))
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, decodedImg)
|
||||
})
|
||||
|
||||
t.Run("strip PNG metadata", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// Encode as PNG first
|
||||
var pngBuf bytes.Buffer
|
||||
err := imaging.Encode(&pngBuf, img, imaging.PNG)
|
||||
require.NoError(t, err)
|
||||
|
||||
strippedData, err := stripImageExif(pngBuf.Bytes(), "image/png")
|
||||
require.NoError(t, err)
|
||||
assert.NotEmpty(t, strippedData)
|
||||
|
||||
// Verify it's still a valid image
|
||||
decodedImg, err := imaging.Decode(bytes.NewReader(strippedData))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, 100, decodedImg.Bounds().Dx())
|
||||
assert.Equal(t, 100, decodedImg.Bounds().Dy())
|
||||
})
|
||||
|
||||
t.Run("handle WebP format by converting to JPEG", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
// WebP format will be converted to JPEG
|
||||
strippedData, err := stripImageExif(originalData, "image/webp")
|
||||
require.NoError(t, err)
|
||||
assert.NotEmpty(t, strippedData)
|
||||
|
||||
// Verify it's a valid image
|
||||
decodedImg, err := imaging.Decode(bytes.NewReader(strippedData))
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, decodedImg)
|
||||
})
|
||||
|
||||
t.Run("handle HEIC format by converting to JPEG", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
strippedData, err := stripImageExif(originalData, "image/heic")
|
||||
require.NoError(t, err)
|
||||
assert.NotEmpty(t, strippedData)
|
||||
|
||||
// Verify it's a valid image
|
||||
decodedImg, err := imaging.Decode(bytes.NewReader(strippedData))
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, decodedImg)
|
||||
})
|
||||
|
||||
t.Run("return error for invalid image data", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
invalidData := []byte("not an image")
|
||||
_, err := stripImageExif(invalidData, "image/jpeg")
|
||||
assert.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "failed to decode image")
|
||||
})
|
||||
|
||||
t.Run("return error for empty image data", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
emptyData := []byte{}
|
||||
_, err := stripImageExif(emptyData, "image/jpeg")
|
||||
assert.Error(t, err)
|
||||
})
|
||||
}
|
||||
|
||||
func TestValidateImagePixelCountRejectsOversizedDimensions(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
err := validateImagePixelCount(testPNGHeaderWithDimensions(100_000, 100_000))
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "image dimensions exceed maximum")
|
||||
}
|
||||
|
||||
func TestStripImageExifRejectsOversizedDimensionsBeforeDecode(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
_, err := stripImageExif(testPNGHeaderWithDimensions(100_000, 100_000), "image/png")
|
||||
require.Error(t, err)
|
||||
require.Contains(t, err.Error(), "image dimensions exceed maximum")
|
||||
}
|
||||
|
||||
func testPNGHeaderWithDimensions(width, height uint32) []byte {
|
||||
var buf bytes.Buffer
|
||||
buf.Write([]byte{0x89, 'P', 'N', 'G', '\r', '\n', 0x1a, '\n'})
|
||||
|
||||
ihdr := make([]byte, 13)
|
||||
binary.BigEndian.PutUint32(ihdr[0:4], width)
|
||||
binary.BigEndian.PutUint32(ihdr[4:8], height)
|
||||
ihdr[8] = 8
|
||||
ihdr[9] = 2
|
||||
|
||||
writePNGChunk(&buf, "IHDR", ihdr)
|
||||
writePNGChunk(&buf, "IEND", nil)
|
||||
return buf.Bytes()
|
||||
}
|
||||
|
||||
func writePNGChunk(buf *bytes.Buffer, chunkType string, data []byte) {
|
||||
_ = binary.Write(buf, binary.BigEndian, uint32(len(data)))
|
||||
buf.WriteString(chunkType)
|
||||
buf.Write(data)
|
||||
crc := crc32.ChecksumIEEE(append([]byte(chunkType), data...))
|
||||
_ = binary.Write(buf, binary.BigEndian, crc)
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func convertMotionMediaFromStore(motion *storepb.MotionMedia) *v1pb.MotionMedia {
|
||||
if motion == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
return &v1pb.MotionMedia{
|
||||
Family: v1pb.MotionMediaFamily(motion.Family),
|
||||
Role: v1pb.MotionMediaRole(motion.Role),
|
||||
GroupId: motion.GroupId,
|
||||
PresentationTimestampUs: motion.PresentationTimestampUs,
|
||||
HasEmbeddedVideo: motion.HasEmbeddedVideo,
|
||||
}
|
||||
}
|
||||
|
||||
func convertMotionMediaToStore(motion *v1pb.MotionMedia) *storepb.MotionMedia {
|
||||
if motion == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
return &storepb.MotionMedia{
|
||||
Family: storepb.MotionMediaFamily(motion.Family),
|
||||
Role: storepb.MotionMediaRole(motion.Role),
|
||||
GroupId: motion.GroupId,
|
||||
PresentationTimestampUs: motion.PresentationTimestampUs,
|
||||
HasEmbeddedVideo: motion.HasEmbeddedVideo,
|
||||
}
|
||||
}
|
||||
|
||||
func getAttachmentMotionMedia(attachment *store.Attachment) *storepb.MotionMedia {
|
||||
if attachment == nil || attachment.Payload == nil {
|
||||
return nil
|
||||
}
|
||||
return attachment.Payload.MotionMedia
|
||||
}
|
||||
|
||||
func isAndroidMotionContainer(motion *storepb.MotionMedia) bool {
|
||||
return motion != nil &&
|
||||
motion.Family == storepb.MotionMediaFamily_ANDROID_MOTION_PHOTO &&
|
||||
motion.Role == storepb.MotionMediaRole_CONTAINER &&
|
||||
motion.HasEmbeddedVideo
|
||||
}
|
||||
|
||||
func ensureAttachmentPayload(payload *storepb.AttachmentPayload) *storepb.AttachmentPayload {
|
||||
if payload != nil {
|
||||
return payload
|
||||
}
|
||||
return &storepb.AttachmentPayload{}
|
||||
}
|
||||
|
||||
func isMultiMemberMotionGroup(attachments []*store.Attachment) bool {
|
||||
if len(attachments) < 2 {
|
||||
return false
|
||||
}
|
||||
for _, attachment := range attachments {
|
||||
motion := getAttachmentMotionMedia(attachment)
|
||||
if motion == nil || motion.GroupId == "" {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
@@ -0,0 +1,831 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/binary"
|
||||
"fmt"
|
||||
"image"
|
||||
"io"
|
||||
"log/slog"
|
||||
"mime"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/disintegration/imaging"
|
||||
"github.com/pkg/errors"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
"google.golang.org/protobuf/types/known/emptypb"
|
||||
"google.golang.org/protobuf/types/known/timestamppb"
|
||||
|
||||
"github.com/usememos/memos/internal/filter"
|
||||
"github.com/usememos/memos/internal/motionphoto"
|
||||
"github.com/usememos/memos/internal/profile"
|
||||
"github.com/usememos/memos/internal/storage/s3"
|
||||
"github.com/usememos/memos/internal/util"
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
const (
|
||||
// The upload memory buffer is 32 MiB.
|
||||
// It should be kept low, so RAM usage doesn't get out of control.
|
||||
// This is unrelated to maximum upload size limit, which is now set through system setting.
|
||||
MaxUploadBufferSizeBytes = 32 << 20
|
||||
MebiByte = 1024 * 1024
|
||||
// ThumbnailCacheFolder is the folder name where the thumbnail images are stored.
|
||||
ThumbnailCacheFolder = ".thumbnail_cache"
|
||||
|
||||
// defaultJPEGQuality is the JPEG quality used when re-encoding images for EXIF stripping.
|
||||
// Quality 95 maintains visual quality while ensuring metadata is removed.
|
||||
defaultJPEGQuality = 95
|
||||
maxBatchDeleteAttachments = 100
|
||||
maxImagePixels = 50_000_000
|
||||
)
|
||||
|
||||
var SupportedThumbnailMimeTypes = []string{
|
||||
"image/png",
|
||||
"image/jpeg",
|
||||
}
|
||||
|
||||
// exifCapableImageTypes defines image formats that may contain EXIF metadata.
|
||||
// These formats will have their EXIF metadata stripped on upload for privacy.
|
||||
var exifCapableImageTypes = map[string]bool{
|
||||
"image/jpeg": true,
|
||||
"image/jpg": true,
|
||||
"image/tiff": true,
|
||||
"image/webp": true,
|
||||
"image/heic": true,
|
||||
"image/heif": true,
|
||||
}
|
||||
|
||||
func (s *APIV1Service) CreateAttachment(ctx context.Context, request *v1pb.CreateAttachmentRequest) (*v1pb.Attachment, error) {
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
|
||||
// Validate required fields
|
||||
if request.Attachment == nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "attachment is required")
|
||||
}
|
||||
if request.Attachment.Filename == "" {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "filename is required")
|
||||
}
|
||||
if !validateFilename(request.Attachment.Filename) {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "filename contains invalid characters or format")
|
||||
}
|
||||
normalizedMimeType := request.Attachment.Type
|
||||
if normalizedMimeType == "" {
|
||||
ext := filepath.Ext(request.Attachment.Filename)
|
||||
mimeType := mime.TypeByExtension(ext)
|
||||
if mimeType == "" {
|
||||
mimeType = http.DetectContentType(request.Attachment.Content)
|
||||
}
|
||||
if normalizedType, ok := normalizeMimeType(mimeType); ok {
|
||||
normalizedMimeType = normalizedType
|
||||
}
|
||||
}
|
||||
if normalizedMimeType == "" {
|
||||
normalizedMimeType = "application/octet-stream"
|
||||
}
|
||||
normalizedType, ok := normalizeMimeType(normalizedMimeType)
|
||||
if !ok {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid MIME type format")
|
||||
}
|
||||
request.Attachment.Type = normalizedType
|
||||
|
||||
attachmentUID, err := ValidateAndGenerateUID(request.AttachmentId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
create := &store.Attachment{
|
||||
UID: attachmentUID,
|
||||
CreatorID: user.ID,
|
||||
Filename: request.Attachment.Filename,
|
||||
Type: request.Attachment.Type,
|
||||
}
|
||||
|
||||
inputMotionMedia, err := validateClientMotionMedia(request.Attachment.MotionMedia, attachmentUID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if inputMotionMedia != nil {
|
||||
create.Payload = ensureAttachmentPayload(create.Payload)
|
||||
create.Payload.MotionMedia = inputMotionMedia
|
||||
}
|
||||
|
||||
instanceStorageSetting, err := s.Store.GetInstanceStorageSetting(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get instance storage setting: %v", err)
|
||||
}
|
||||
size := binary.Size(request.Attachment.Content)
|
||||
uploadSizeLimit := int(instanceStorageSetting.UploadSizeLimitMb) * MebiByte
|
||||
if uploadSizeLimit == 0 {
|
||||
uploadSizeLimit = MaxUploadBufferSizeBytes
|
||||
}
|
||||
if size > uploadSizeLimit {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "file size exceeds the limit")
|
||||
}
|
||||
create.Size = int64(size)
|
||||
create.Blob = request.Attachment.Content
|
||||
|
||||
if request.Attachment.Memo != nil {
|
||||
memoUID, err := ExtractMemoUIDFromName(*request.Attachment.Memo)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid memo name: %v", err)
|
||||
}
|
||||
memo, err := s.Store.GetMemo(ctx, &store.FindMemo{UID: &memoUID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to find memo: %v", err)
|
||||
}
|
||||
if memo == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "memo not found: %s", *request.Attachment.Memo)
|
||||
}
|
||||
if !canModifyMemo(user, memo) {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
create.MemoID = &memo.ID
|
||||
}
|
||||
|
||||
if create.Payload == nil || create.Payload.MotionMedia == nil {
|
||||
if detectedMotion := detectAndroidMotionMedia(create.Blob, create.Type, attachmentUID); detectedMotion != nil {
|
||||
create.Payload = ensureAttachmentPayload(create.Payload)
|
||||
create.Payload.MotionMedia = detectedMotion
|
||||
}
|
||||
}
|
||||
|
||||
// Strip EXIF metadata from images for privacy protection.
|
||||
// This removes sensitive information like GPS location, device details, etc.
|
||||
if shouldStripExif(create.Type) && !isAndroidMotionContainer(create.Payload.GetMotionMedia()) {
|
||||
release, err := s.acquireImageProcessingSlot(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.ResourceExhausted, "too many image processing requests")
|
||||
}
|
||||
strippedBlob, stripErr := stripImageExif(create.Blob, create.Type)
|
||||
release()
|
||||
if stripErr != nil {
|
||||
// Log warning but continue with original image to ensure uploads don't fail.
|
||||
slog.Warn("failed to strip EXIF metadata from image",
|
||||
slog.String("type", create.Type),
|
||||
slog.String("filename", create.Filename),
|
||||
slog.String("error", stripErr.Error()))
|
||||
} else {
|
||||
create.Blob = strippedBlob
|
||||
create.Size = int64(len(strippedBlob))
|
||||
}
|
||||
}
|
||||
|
||||
if err := SaveAttachmentBlob(ctx, s.Profile, s.Store, create); err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to save attachment blob: %v", err)
|
||||
}
|
||||
|
||||
attachment, err := s.Store.CreateAttachment(ctx, create)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to create attachment: %v", err)
|
||||
}
|
||||
|
||||
return convertAttachmentFromStore(attachment), nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) ListAttachments(ctx context.Context, request *v1pb.ListAttachmentsRequest) (*v1pb.ListAttachmentsResponse, error) {
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
|
||||
// Set default page size
|
||||
pageSize := int(request.PageSize)
|
||||
if pageSize <= 0 {
|
||||
pageSize = 50
|
||||
}
|
||||
if pageSize > 1000 {
|
||||
pageSize = 1000
|
||||
}
|
||||
|
||||
// Parse page token for offset
|
||||
offset := 0
|
||||
if request.PageToken != "" {
|
||||
// Simple implementation: page token is the offset as string
|
||||
// In production, you might want to use encrypted tokens
|
||||
if parsed, err := fmt.Sscanf(request.PageToken, "%d", &offset); err != nil || parsed != 1 {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid page token")
|
||||
}
|
||||
}
|
||||
|
||||
findAttachment := &store.FindAttachment{
|
||||
CreatorID: &user.ID,
|
||||
Limit: &pageSize,
|
||||
Offset: &offset,
|
||||
}
|
||||
|
||||
// Parse filter if provided
|
||||
if request.Filter != "" {
|
||||
if err := s.validateAttachmentFilter(ctx, request.Filter); err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid filter: %v", err)
|
||||
}
|
||||
findAttachment.Filters = append(findAttachment.Filters, request.Filter)
|
||||
}
|
||||
|
||||
attachments, err := s.Store.ListAttachments(ctx, findAttachment)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to list attachments: %v", err)
|
||||
}
|
||||
|
||||
response := &v1pb.ListAttachmentsResponse{}
|
||||
|
||||
for _, attachment := range attachments {
|
||||
response.Attachments = append(response.Attachments, convertAttachmentFromStore(attachment))
|
||||
}
|
||||
|
||||
// For simplicity, set total size to the number of returned attachments.
|
||||
// In a full implementation, you'd want a separate count query
|
||||
response.TotalSize = int32(len(response.Attachments))
|
||||
|
||||
// Set next page token if we got the full page size (indicating there might be more)
|
||||
if len(attachments) == pageSize {
|
||||
response.NextPageToken = fmt.Sprintf("%d", offset+pageSize)
|
||||
}
|
||||
|
||||
return response, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) GetAttachment(ctx context.Context, request *v1pb.GetAttachmentRequest) (*v1pb.Attachment, error) {
|
||||
attachmentUID, err := ExtractAttachmentUIDFromName(request.Name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid attachment id: %v", err)
|
||||
}
|
||||
attachment, err := s.Store.GetAttachment(ctx, &store.FindAttachment{UID: &attachmentUID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get attachment: %v", err)
|
||||
}
|
||||
if attachment == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "attachment not found")
|
||||
}
|
||||
|
||||
// Check access permission based on linked memo visibility.
|
||||
if err := s.checkAttachmentAccess(ctx, attachment); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return convertAttachmentFromStore(attachment), nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) UpdateAttachment(ctx context.Context, request *v1pb.UpdateAttachmentRequest) (*v1pb.Attachment, error) {
|
||||
attachmentUID, err := ExtractAttachmentUIDFromName(request.Attachment.Name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid attachment id: %v", err)
|
||||
}
|
||||
if request.UpdateMask == nil || len(request.UpdateMask.Paths) == 0 {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "update mask is required")
|
||||
}
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
attachment, err := s.Store.GetAttachment(ctx, &store.FindAttachment{UID: &attachmentUID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get attachment: %v", err)
|
||||
}
|
||||
if attachment == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "attachment not found")
|
||||
}
|
||||
// Only the creator or admin can update the attachment.
|
||||
if attachment.CreatorID != user.ID && !isSuperUser(user) {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
|
||||
currentTs := time.Now().Unix()
|
||||
update := &store.UpdateAttachment{
|
||||
ID: attachment.ID,
|
||||
UpdatedTs: ¤tTs,
|
||||
}
|
||||
for _, field := range request.UpdateMask.Paths {
|
||||
if field == "filename" {
|
||||
if !validateFilename(request.Attachment.Filename) {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "filename contains invalid characters or format")
|
||||
}
|
||||
update.Filename = &request.Attachment.Filename
|
||||
}
|
||||
}
|
||||
|
||||
if err := s.Store.UpdateAttachment(ctx, update); err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to update attachment: %v", err)
|
||||
}
|
||||
return s.GetAttachment(ctx, &v1pb.GetAttachmentRequest{
|
||||
Name: request.Attachment.Name,
|
||||
})
|
||||
}
|
||||
|
||||
func (s *APIV1Service) DeleteAttachment(ctx context.Context, request *v1pb.DeleteAttachmentRequest) (*emptypb.Empty, error) {
|
||||
attachmentUID, err := ExtractAttachmentUIDFromName(request.Name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid attachment id: %v", err)
|
||||
}
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
attachment, err := s.Store.GetAttachment(ctx, &store.FindAttachment{
|
||||
UID: &attachmentUID,
|
||||
CreatorID: &user.ID,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to find attachment: %v", err)
|
||||
}
|
||||
if attachment == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "attachment not found")
|
||||
}
|
||||
// Delete the attachment from the database.
|
||||
if err := s.Store.DeleteAttachment(ctx, &store.DeleteAttachment{
|
||||
ID: attachment.ID,
|
||||
}); err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to delete attachment: %v", err)
|
||||
}
|
||||
return &emptypb.Empty{}, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) BatchDeleteAttachments(ctx context.Context, request *v1pb.BatchDeleteAttachmentsRequest) (*emptypb.Empty, error) {
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
if len(request.Names) == 0 {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "attachment names are required")
|
||||
}
|
||||
if len(request.Names) > maxBatchDeleteAttachments {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "too many attachment names; max %d", maxBatchDeleteAttachments)
|
||||
}
|
||||
|
||||
attachments := make([]*store.Attachment, 0, len(request.Names))
|
||||
seen := make(map[string]bool, len(request.Names))
|
||||
for _, name := range request.Names {
|
||||
if name == "" {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "attachment name is required")
|
||||
}
|
||||
if seen[name] {
|
||||
continue
|
||||
}
|
||||
seen[name] = true
|
||||
|
||||
attachmentUID, err := ExtractAttachmentUIDFromName(name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid attachment id: %v", err)
|
||||
}
|
||||
attachment, err := s.Store.GetAttachment(ctx, &store.FindAttachment{UID: &attachmentUID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get attachment: %v", err)
|
||||
}
|
||||
if attachment == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "attachment not found")
|
||||
}
|
||||
if attachment.CreatorID != user.ID && !isSuperUser(user) {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
attachments = append(attachments, attachment)
|
||||
}
|
||||
|
||||
if err := s.Store.DeleteAttachments(ctx, attachments); err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to delete attachments: %v", err)
|
||||
}
|
||||
|
||||
return &emptypb.Empty{}, nil
|
||||
}
|
||||
|
||||
func convertAttachmentFromStore(attachment *store.Attachment) *v1pb.Attachment {
|
||||
attachmentMessage := &v1pb.Attachment{
|
||||
Name: fmt.Sprintf("%s%s", AttachmentNamePrefix, attachment.UID),
|
||||
CreateTime: timestamppb.New(time.Unix(attachment.CreatedTs, 0)),
|
||||
Filename: attachment.Filename,
|
||||
Type: attachment.Type,
|
||||
Size: attachment.Size,
|
||||
MotionMedia: convertMotionMediaFromStore(getAttachmentMotionMedia(attachment)),
|
||||
}
|
||||
if attachment.MemoUID != nil && *attachment.MemoUID != "" {
|
||||
memoName := fmt.Sprintf("%s%s", MemoNamePrefix, *attachment.MemoUID)
|
||||
attachmentMessage.Memo = &memoName
|
||||
}
|
||||
if attachment.StorageType == storepb.AttachmentStorageType_EXTERNAL || attachment.StorageType == storepb.AttachmentStorageType_S3 {
|
||||
attachmentMessage.ExternalLink = attachment.Reference
|
||||
}
|
||||
|
||||
return attachmentMessage
|
||||
}
|
||||
|
||||
// SaveAttachmentBlob saves the blob of attachment based on the storage config.
|
||||
func SaveAttachmentBlob(ctx context.Context, profile *profile.Profile, stores *store.Store, create *store.Attachment) error {
|
||||
instanceStorageSetting, err := stores.GetInstanceStorageSetting(ctx)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "Failed to find instance storage setting")
|
||||
}
|
||||
|
||||
if instanceStorageSetting.StorageType == storepb.InstanceStorageSetting_LOCAL {
|
||||
filepathTemplate := "assets/{timestamp}_{uuid}_{filename}"
|
||||
if instanceStorageSetting.FilepathTemplate != "" {
|
||||
filepathTemplate = instanceStorageSetting.FilepathTemplate
|
||||
}
|
||||
|
||||
internalPath := filepathTemplate
|
||||
if !strings.Contains(internalPath, "{filename}") {
|
||||
internalPath = filepath.Join(internalPath, "{filename}")
|
||||
}
|
||||
internalPath = replaceFilenameWithPathTemplate(internalPath, create.Filename)
|
||||
internalPath = filepath.ToSlash(internalPath)
|
||||
|
||||
// Ensure the directory exists.
|
||||
osPath := filepath.FromSlash(internalPath)
|
||||
if !filepath.IsAbs(osPath) {
|
||||
osPath = filepath.Join(profile.Data, osPath)
|
||||
}
|
||||
osPath = ensureUniqueLocalAttachmentPath(osPath, create.UID)
|
||||
internalPath = filepath.ToSlash(osPath)
|
||||
if !filepath.IsAbs(filepath.FromSlash(internalPath)) {
|
||||
internalPath, err = filepath.Rel(profile.Data, osPath)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "Failed to get relative path")
|
||||
}
|
||||
internalPath = filepath.ToSlash(internalPath)
|
||||
}
|
||||
dir := filepath.Dir(osPath)
|
||||
if err = os.MkdirAll(dir, os.ModePerm); err != nil {
|
||||
return errors.Wrap(err, "Failed to create directory")
|
||||
}
|
||||
|
||||
// Write the blob to the file.
|
||||
if err := os.WriteFile(osPath, create.Blob, 0644); err != nil {
|
||||
return errors.Wrap(err, "Failed to write file")
|
||||
}
|
||||
create.Reference = internalPath
|
||||
create.Blob = nil
|
||||
create.StorageType = storepb.AttachmentStorageType_LOCAL
|
||||
} else if instanceStorageSetting.StorageType == storepb.InstanceStorageSetting_S3 {
|
||||
s3Config := instanceStorageSetting.S3Config
|
||||
if s3Config == nil {
|
||||
return errors.Errorf("No activated external storage found")
|
||||
}
|
||||
s3Client, err := s3.NewClient(ctx, s3Config)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "Failed to create s3 client")
|
||||
}
|
||||
|
||||
filepathTemplate := instanceStorageSetting.FilepathTemplate
|
||||
if !strings.Contains(filepathTemplate, "{filename}") {
|
||||
filepathTemplate = filepath.Join(filepathTemplate, "{filename}")
|
||||
}
|
||||
filepathTemplate = replaceFilenameWithPathTemplate(filepathTemplate, create.Filename)
|
||||
key, err := s3Client.UploadObject(ctx, filepathTemplate, create.Type, bytes.NewReader(create.Blob))
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "Failed to upload via s3 client")
|
||||
}
|
||||
presignURL, err := s3Client.PresignGetObject(ctx, key)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "Failed to presign via s3 client")
|
||||
}
|
||||
|
||||
create.Reference = presignURL
|
||||
create.Blob = nil
|
||||
create.StorageType = storepb.AttachmentStorageType_S3
|
||||
payload := ensureAttachmentPayload(create.Payload)
|
||||
payload.Payload = &storepb.AttachmentPayload_S3Object_{
|
||||
S3Object: &storepb.AttachmentPayload_S3Object{
|
||||
S3Config: s3Config,
|
||||
Key: key,
|
||||
LastPresignedTime: timestamppb.New(time.Now()),
|
||||
},
|
||||
}
|
||||
create.Payload = payload
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) GetAttachmentBlob(attachment *store.Attachment) ([]byte, error) {
|
||||
// For local storage, read the file from the local disk.
|
||||
if attachment.StorageType == storepb.AttachmentStorageType_LOCAL {
|
||||
attachmentPath := filepath.FromSlash(attachment.Reference)
|
||||
if !filepath.IsAbs(attachmentPath) {
|
||||
attachmentPath = filepath.Join(s.Profile.Data, attachmentPath)
|
||||
}
|
||||
|
||||
file, err := os.Open(attachmentPath)
|
||||
if err != nil {
|
||||
if os.IsNotExist(err) {
|
||||
return nil, errors.Wrap(err, "file not found")
|
||||
}
|
||||
return nil, errors.Wrap(err, "failed to open the file")
|
||||
}
|
||||
defer file.Close()
|
||||
blob, err := io.ReadAll(file)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to read the file")
|
||||
}
|
||||
return blob, nil
|
||||
}
|
||||
// For S3 storage, download the file from S3.
|
||||
if attachment.StorageType == storepb.AttachmentStorageType_S3 {
|
||||
if attachment.Payload == nil {
|
||||
return nil, errors.New("attachment payload is missing")
|
||||
}
|
||||
s3Object := attachment.Payload.GetS3Object()
|
||||
if s3Object == nil {
|
||||
return nil, errors.New("S3 object payload is missing")
|
||||
}
|
||||
if s3Object.S3Config == nil {
|
||||
return nil, errors.New("S3 config is missing")
|
||||
}
|
||||
if s3Object.Key == "" {
|
||||
return nil, errors.New("S3 object key is missing")
|
||||
}
|
||||
|
||||
s3Client, err := s3.NewClient(context.Background(), s3Object.S3Config)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to create S3 client")
|
||||
}
|
||||
|
||||
blob, err := s3Client.GetObject(context.Background(), s3Object.Key)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to get object from S3")
|
||||
}
|
||||
return blob, nil
|
||||
}
|
||||
// For database storage, return the blob from the database.
|
||||
return attachment.Blob, nil
|
||||
}
|
||||
|
||||
var fileKeyPattern = regexp.MustCompile(`\{[a-z]{1,9}\}`)
|
||||
|
||||
func replaceFilenameWithPathTemplate(path, filename string) string {
|
||||
t := time.Now()
|
||||
path = fileKeyPattern.ReplaceAllStringFunc(path, func(s string) string {
|
||||
switch s {
|
||||
case "{filename}":
|
||||
return filename
|
||||
case "{timestamp}":
|
||||
return fmt.Sprintf("%d", t.Unix())
|
||||
case "{year}":
|
||||
return fmt.Sprintf("%d", t.Year())
|
||||
case "{month}":
|
||||
return fmt.Sprintf("%02d", t.Month())
|
||||
case "{day}":
|
||||
return fmt.Sprintf("%02d", t.Day())
|
||||
case "{hour}":
|
||||
return fmt.Sprintf("%02d", t.Hour())
|
||||
case "{minute}":
|
||||
return fmt.Sprintf("%02d", t.Minute())
|
||||
case "{second}":
|
||||
return fmt.Sprintf("%02d", t.Second())
|
||||
case "{uuid}":
|
||||
return util.GenUUID()
|
||||
default:
|
||||
return s
|
||||
}
|
||||
})
|
||||
return path
|
||||
}
|
||||
|
||||
func ensureUniqueLocalAttachmentPath(path, uid string) string {
|
||||
if _, err := os.Stat(path); err != nil {
|
||||
return path
|
||||
}
|
||||
|
||||
ext := filepath.Ext(path)
|
||||
base := strings.TrimSuffix(path, ext)
|
||||
return base + "_" + uid + ext
|
||||
}
|
||||
|
||||
func validateFilename(filename string) bool {
|
||||
// Reject path traversal attempts and make sure no additional directories are created
|
||||
if !filepath.IsLocal(filename) || strings.ContainsAny(filename, "/\\") {
|
||||
return false
|
||||
}
|
||||
|
||||
// Reject filenames starting or ending with spaces or periods
|
||||
if strings.HasPrefix(filename, " ") || strings.HasSuffix(filename, " ") ||
|
||||
strings.HasPrefix(filename, ".") || strings.HasSuffix(filename, ".") {
|
||||
return false
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
func normalizeMimeType(mimeType string) (string, bool) {
|
||||
mimeType = strings.TrimSpace(mimeType)
|
||||
if mimeType == "" || len(mimeType) > 255 {
|
||||
return "", false
|
||||
}
|
||||
|
||||
mediaType, _, err := mime.ParseMediaType(mimeType)
|
||||
if err != nil || mediaType == "" || len(mediaType) > 255 {
|
||||
return "", false
|
||||
}
|
||||
|
||||
return mediaType, true
|
||||
}
|
||||
|
||||
func (s *APIV1Service) validateAttachmentFilter(ctx context.Context, filterStr string) error {
|
||||
if filterStr == "" {
|
||||
return errors.New("filter cannot be empty")
|
||||
}
|
||||
|
||||
engine, err := filter.DefaultAttachmentEngine()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var dialect filter.DialectName
|
||||
switch s.Profile.Driver {
|
||||
case "mysql":
|
||||
dialect = filter.DialectMySQL
|
||||
case "postgres":
|
||||
dialect = filter.DialectPostgres
|
||||
default:
|
||||
dialect = filter.DialectSQLite
|
||||
}
|
||||
|
||||
if _, err := engine.CompileToStatement(ctx, filterStr, filter.RenderOptions{Dialect: dialect}); err != nil {
|
||||
return errors.Wrap(err, "failed to compile filter")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// checkAttachmentAccess verifies the user has permission to access the attachment.
|
||||
// For unlinked attachments (no memo), only the creator can access.
|
||||
// For linked attachments, access follows the memo's visibility rules.
|
||||
func (s *APIV1Service) checkAttachmentAccess(ctx context.Context, attachment *store.Attachment) error {
|
||||
user, _ := s.fetchCurrentUser(ctx)
|
||||
|
||||
// For unlinked attachments, only the creator can access.
|
||||
if attachment.MemoID == nil {
|
||||
if user == nil {
|
||||
return status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
if attachment.CreatorID != user.ID && !isSuperUser(user) {
|
||||
return status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// For linked attachments, check memo visibility.
|
||||
memo, err := s.Store.GetMemo(ctx, &store.FindMemo{ID: attachment.MemoID})
|
||||
if err != nil {
|
||||
return status.Errorf(codes.Internal, "failed to get memo: %v", err)
|
||||
}
|
||||
if memo == nil {
|
||||
return status.Errorf(codes.NotFound, "memo not found")
|
||||
}
|
||||
|
||||
if memo.Visibility == store.Public {
|
||||
return nil
|
||||
}
|
||||
if user == nil {
|
||||
return status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
if memo.Visibility == store.Private && memo.CreatorID != user.ID && !isSuperUser(user) {
|
||||
return status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateClientMotionMedia(motion *v1pb.MotionMedia, attachmentUID string) (*storepb.MotionMedia, error) {
|
||||
if motion == nil {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
if motion.Family != v1pb.MotionMediaFamily_APPLE_LIVE_PHOTO {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "only Apple Live Photo motion metadata can be provided by clients")
|
||||
}
|
||||
if motion.Role != v1pb.MotionMediaRole_STILL && motion.Role != v1pb.MotionMediaRole_VIDEO {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid Apple Live Photo motion role")
|
||||
}
|
||||
|
||||
storeMotion := convertMotionMediaToStore(motion)
|
||||
if storeMotion.GroupId == "" {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "motion media group_id is required")
|
||||
}
|
||||
if storeMotion.Family == storepb.MotionMediaFamily_ANDROID_MOTION_PHOTO && storeMotion.GroupId == "" {
|
||||
storeMotion.GroupId = attachmentUID
|
||||
}
|
||||
|
||||
return storeMotion, nil
|
||||
}
|
||||
|
||||
func detectAndroidMotionMedia(blob []byte, mimeType, attachmentUID string) *storepb.MotionMedia {
|
||||
if mimeType != "image/jpeg" && mimeType != "image/jpg" {
|
||||
return nil
|
||||
}
|
||||
|
||||
detection := motionphoto.DetectJPEG(blob)
|
||||
if detection == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
return &storepb.MotionMedia{
|
||||
Family: storepb.MotionMediaFamily_ANDROID_MOTION_PHOTO,
|
||||
Role: storepb.MotionMediaRole_CONTAINER,
|
||||
GroupId: attachmentUID,
|
||||
PresentationTimestampUs: detection.PresentationTimestampUs,
|
||||
HasEmbeddedVideo: true,
|
||||
}
|
||||
}
|
||||
|
||||
// shouldStripExif checks if the MIME type is an image format that may contain EXIF metadata.
|
||||
// Returns true for formats like JPEG, TIFF, WebP, HEIC, and HEIF which commonly contain
|
||||
// privacy-sensitive metadata such as GPS coordinates, camera settings, and device information.
|
||||
func shouldStripExif(mimeType string) bool {
|
||||
return exifCapableImageTypes[mimeType]
|
||||
}
|
||||
|
||||
func (s *APIV1Service) acquireImageProcessingSlot(ctx context.Context) (func(), error) {
|
||||
if s.imageProcessingSemaphore == nil {
|
||||
return func() {}, nil
|
||||
}
|
||||
if err := s.imageProcessingSemaphore.Acquire(ctx, 1); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return func() {
|
||||
s.imageProcessingSemaphore.Release(1)
|
||||
}, nil
|
||||
}
|
||||
|
||||
func validateImagePixelCount(imageData []byte) error {
|
||||
config, _, err := image.DecodeConfig(bytes.NewReader(imageData))
|
||||
if err != nil {
|
||||
// Some formats supported by imaging do not expose dimensions through
|
||||
// the standard image registry. Let the full decoder handle those.
|
||||
return nil //nolint:nilerr
|
||||
}
|
||||
if config.Width <= 0 || config.Height <= 0 {
|
||||
return errors.New("invalid image dimensions")
|
||||
}
|
||||
if config.Width > maxImagePixels/config.Height {
|
||||
return errors.Errorf("image dimensions exceed maximum of %d pixels", maxImagePixels)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// stripImageExif removes EXIF metadata from image files by decoding and re-encoding them.
|
||||
// This prevents exposure of sensitive metadata such as GPS location, camera details, and timestamps.
|
||||
//
|
||||
// The function preserves the correct image orientation by applying EXIF orientation tags
|
||||
// during decoding before stripping all metadata. Images are re-encoded with high quality
|
||||
// to minimize visual degradation.
|
||||
//
|
||||
// Supported formats:
|
||||
// - JPEG/JPG: Re-encoded as JPEG with quality 95
|
||||
// - PNG: Re-encoded as PNG (lossless)
|
||||
// - TIFF/WebP/HEIC/HEIF: Re-encoded as JPEG with quality 95
|
||||
//
|
||||
// Returns the cleaned image data without any EXIF metadata, or an error if processing fails.
|
||||
func stripImageExif(imageData []byte, mimeType string) ([]byte, error) {
|
||||
if err := validateImagePixelCount(imageData); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Decode image with automatic EXIF orientation correction.
|
||||
// This ensures the image displays correctly after metadata removal.
|
||||
img, err := imaging.Decode(bytes.NewReader(imageData), imaging.AutoOrientation(true))
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to decode image")
|
||||
}
|
||||
|
||||
// Re-encode the image without EXIF metadata.
|
||||
var buf bytes.Buffer
|
||||
var encodeErr error
|
||||
|
||||
if mimeType == "image/png" {
|
||||
// Preserve PNG format for lossless encoding
|
||||
encodeErr = imaging.Encode(&buf, img, imaging.PNG)
|
||||
} else {
|
||||
// For JPEG, TIFF, WebP, HEIC, HEIF - re-encode as JPEG.
|
||||
// This ensures EXIF is stripped and provides good compression.
|
||||
encodeErr = imaging.Encode(&buf, img, imaging.JPEG, imaging.JPEGQuality(defaultJPEGQuality))
|
||||
}
|
||||
|
||||
if encodeErr != nil {
|
||||
return nil, errors.Wrap(encodeErr, "failed to encode image")
|
||||
}
|
||||
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
@@ -0,0 +1,796 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/metadata"
|
||||
"google.golang.org/grpc/status"
|
||||
"google.golang.org/protobuf/types/known/emptypb"
|
||||
"google.golang.org/protobuf/types/known/timestamppb"
|
||||
|
||||
"github.com/usememos/memos/internal/idp"
|
||||
"github.com/usememos/memos/internal/idp/oauth2"
|
||||
"github.com/usememos/memos/internal/util"
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/server/auth"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
const (
|
||||
unmatchedUsernameAndPasswordError = "unmatched username and password"
|
||||
)
|
||||
|
||||
// GetCurrentUser returns the authenticated user's information.
|
||||
// Validates the access token and returns user details.
|
||||
//
|
||||
// Authentication: Required (access token).
|
||||
// Returns: User information.
|
||||
func (s *APIV1Service) GetCurrentUser(ctx context.Context, _ *v1pb.GetCurrentUserRequest) (*v1pb.GetCurrentUserResponse, error) {
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "failed to get current user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
// Clear auth cookies
|
||||
if err := s.clearAuthCookies(ctx); err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to clear auth cookies: %v", err)
|
||||
}
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not found")
|
||||
}
|
||||
|
||||
return &v1pb.GetCurrentUserResponse{
|
||||
User: convertUserFromStore(user, user),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// SignIn authenticates a user with credentials and returns tokens.
|
||||
// On success, returns an access token and sets a refresh token cookie.
|
||||
//
|
||||
// Supports two authentication methods:
|
||||
// 1. Password-based authentication (username + password).
|
||||
// 2. SSO authentication (OAuth2 authorization code).
|
||||
//
|
||||
// Authentication: Not required (public endpoint).
|
||||
// Returns: User info, access token, and token expiry.
|
||||
func (s *APIV1Service) SignIn(ctx context.Context, request *v1pb.SignInRequest) (*v1pb.SignInResponse, error) {
|
||||
var existingUser *store.User
|
||||
|
||||
// Authentication Method 1: Password-based authentication
|
||||
if passwordCredentials := request.GetPasswordCredentials(); passwordCredentials != nil {
|
||||
user, err := s.Store.GetUser(ctx, &store.FindUser{
|
||||
Username: &passwordCredentials.Username,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get user, error: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, unmatchedUsernameAndPasswordError)
|
||||
}
|
||||
// Compare the stored hashed password, with the hashed version of the password that was received.
|
||||
if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(passwordCredentials.Password)); err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, unmatchedUsernameAndPasswordError)
|
||||
}
|
||||
instanceGeneralSetting, err := s.Store.GetInstanceGeneralSetting(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get instance general setting, error: %v", err)
|
||||
}
|
||||
// Check if the password auth in is allowed.
|
||||
if instanceGeneralSetting.DisallowPasswordAuth && user.Role == store.RoleUser {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "password signin is not allowed")
|
||||
}
|
||||
existingUser = user
|
||||
} else if ssoCredentials := request.GetSsoCredentials(); ssoCredentials != nil {
|
||||
// Authentication Method 2: SSO (OAuth2) authentication
|
||||
identityProvider, userInfo, err := s.resolveSSOIdentity(ctx, ssoCredentials.IdpName, ssoCredentials.Code, ssoCredentials.RedirectUri, ssoCredentials.CodeVerifier)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
user, err := s.resolveSSOUser(ctx, nil, identityProvider, userInfo)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
existingUser = user
|
||||
}
|
||||
|
||||
if existingUser == nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid credentials")
|
||||
}
|
||||
if existingUser.RowStatus == store.Archived {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "user has been archived with username %s", existingUser.Username)
|
||||
}
|
||||
|
||||
accessToken, accessExpiresAt, err := s.doSignIn(ctx, existingUser)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to sign in: %v", err)
|
||||
}
|
||||
|
||||
return &v1pb.SignInResponse{
|
||||
User: convertUserFromStore(existingUser, existingUser),
|
||||
AccessToken: accessToken,
|
||||
AccessTokenExpiresAt: timestamppb.New(accessExpiresAt),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// resolveSSOUser resolves a local user from an external-identity subject, creating the
|
||||
// linkage record (and a new local user if necessary) when first login is allowed.
|
||||
//
|
||||
// Lookup goes through the user_identity table so that userInfo.Identifier is never used
|
||||
// as the local username key. On the miss path, a local user is created with a
|
||||
// UUID-based local username (see deriveSSOUsername) and the (provider, extern_uid)
|
||||
// linkage is inserted in the same flow. When currentUser is provided by a caller
|
||||
// outside AuthService.SignIn, the lookup miss path binds the external identity to
|
||||
// that existing user instead. If the linkage insert loses a race on the unique
|
||||
// (provider, extern_uid) constraint, the winning linkage's user is loaded and
|
||||
// checked against the current user.
|
||||
func (s *APIV1Service) resolveSSOUser(ctx context.Context, currentUser *store.User, identityProvider *storepb.IdentityProvider, userInfo *idp.IdentityProviderUserInfo) (*store.User, error) {
|
||||
provider := identityProvider.Uid
|
||||
externUID := userInfo.Identifier
|
||||
|
||||
user, err := s.getLinkedSSOUser(ctx, provider, externUID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if user != nil {
|
||||
if currentUser != nil && currentUser.ID != user.ID {
|
||||
return nil, status.Errorf(codes.AlreadyExists, "identity provider account is already linked to another user")
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
|
||||
if currentUser != nil {
|
||||
return s.bindSSOIdentityToUser(ctx, currentUser, provider, externUID)
|
||||
}
|
||||
|
||||
// Miss path: enforce the registration gate before creating anything.
|
||||
instanceGeneralSetting, err := s.Store.GetInstanceGeneralSetting(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get instance general setting, error: %v", err)
|
||||
}
|
||||
if instanceGeneralSetting.DisallowUserRegistration {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "user registration is not allowed")
|
||||
}
|
||||
|
||||
password, err := util.RandomString(20)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to generate random password, error: %v", err)
|
||||
}
|
||||
passwordHash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to generate password hash, error: %v", err)
|
||||
}
|
||||
username, err := deriveSSOUsername()
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to derive username, error: %v", err)
|
||||
}
|
||||
user, err = s.Store.CreateUser(ctx, &store.User{
|
||||
Username: username,
|
||||
Role: store.RoleUser,
|
||||
Nickname: userInfo.DisplayName,
|
||||
Email: userInfo.Email,
|
||||
AvatarURL: userInfo.AvatarURL,
|
||||
PasswordHash: string(passwordHash),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to create user, error: %v", err)
|
||||
}
|
||||
|
||||
if _, err := s.Store.CreateUserIdentity(ctx, &store.UserIdentity{
|
||||
UserID: user.ID,
|
||||
Provider: provider,
|
||||
ExternUID: externUID,
|
||||
}); err != nil {
|
||||
// Best-effort cleanup: the provisional user row has no linkage and should not remain.
|
||||
_, _ = s.Store.DeleteUser(ctx, &store.DeleteUser{ID: user.ID})
|
||||
if isUniqueConstraintViolation(err) {
|
||||
// Concurrent first login won the race; load the winning linkage's user.
|
||||
winner, getErr := s.Store.GetUserIdentity(ctx, &store.FindUserIdentity{
|
||||
Provider: &provider,
|
||||
ExternUID: &externUID,
|
||||
})
|
||||
if getErr != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to reload user identity after race, error: %v", getErr)
|
||||
}
|
||||
if winner == nil {
|
||||
return nil, status.Errorf(codes.Internal, "user identity conflict reported but no winning row found")
|
||||
}
|
||||
winnerUser, getErr := s.Store.GetUser(ctx, &store.FindUser{ID: &winner.UserID})
|
||||
if getErr != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get user after race, error: %v", getErr)
|
||||
}
|
||||
if winnerUser == nil {
|
||||
return nil, status.Errorf(codes.Internal, "linked user %d not found after race", winner.UserID)
|
||||
}
|
||||
return winnerUser, nil
|
||||
}
|
||||
return nil, status.Errorf(codes.Internal, "failed to create user identity, error: %v", err)
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) resolveSSOIdentity(ctx context.Context, idpName, code, redirectURI, codeVerifier string) (*storepb.IdentityProvider, *idp.IdentityProviderUserInfo, error) {
|
||||
idpUID, err := ExtractIdentityProviderUIDFromName(idpName)
|
||||
if err != nil {
|
||||
return nil, nil, status.Errorf(codes.InvalidArgument, "invalid identity provider name: %v", err)
|
||||
}
|
||||
identityProvider, err := s.Store.GetIdentityProvider(ctx, &store.FindIdentityProvider{
|
||||
UID: &idpUID,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, nil, status.Errorf(codes.Internal, "failed to get identity provider, error: %v", err)
|
||||
}
|
||||
if identityProvider == nil {
|
||||
return nil, nil, status.Errorf(codes.InvalidArgument, "identity provider not found")
|
||||
}
|
||||
|
||||
var userInfo *idp.IdentityProviderUserInfo
|
||||
if identityProvider.Type == storepb.IdentityProvider_OAUTH2 {
|
||||
oauth2IdentityProvider, err := oauth2.NewIdentityProvider(identityProvider.Config.GetOauth2Config())
|
||||
if err != nil {
|
||||
return nil, nil, status.Errorf(codes.Internal, "failed to create oauth2 identity provider, error: %v", err)
|
||||
}
|
||||
// Pass code_verifier for PKCE support (empty string if not provided for backward compatibility)
|
||||
token, err := oauth2IdentityProvider.ExchangeToken(ctx, redirectURI, code, codeVerifier)
|
||||
if err != nil {
|
||||
return nil, nil, status.Errorf(codes.Internal, "failed to exchange token, error: %v", err)
|
||||
}
|
||||
userInfo, err = oauth2IdentityProvider.UserInfo(ctx, token)
|
||||
if err != nil {
|
||||
return nil, nil, status.Errorf(codes.Internal, "failed to get user info, error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
identifierFilter := identityProvider.IdentifierFilter
|
||||
if identifierFilter != "" {
|
||||
identifierFilterRegex, err := regexp.Compile(identifierFilter)
|
||||
if err != nil {
|
||||
return nil, nil, status.Errorf(codes.Internal, "failed to compile identifier filter regex, error: %v", err)
|
||||
}
|
||||
if !identifierFilterRegex.MatchString(userInfo.Identifier) {
|
||||
return nil, nil, status.Errorf(codes.PermissionDenied, "identifier %s is not allowed", userInfo.Identifier)
|
||||
}
|
||||
}
|
||||
|
||||
return identityProvider, userInfo, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) getLinkedSSOUser(ctx context.Context, provider, externUID string) (*store.User, error) {
|
||||
identity, err := s.Store.GetUserIdentity(ctx, &store.FindUserIdentity{
|
||||
Provider: &provider,
|
||||
ExternUID: &externUID,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get user identity, error: %v", err)
|
||||
}
|
||||
if identity == nil {
|
||||
return nil, nil
|
||||
}
|
||||
user, err := s.Store.GetUser(ctx, &store.FindUser{ID: &identity.UserID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get user, error: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Internal, "linked user %d not found for identity %d", identity.UserID, identity.ID)
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) bindSSOIdentityToUser(ctx context.Context, currentUser *store.User, provider, externUID string) (*store.User, error) {
|
||||
existingForProvider, err := s.Store.GetUserIdentity(ctx, &store.FindUserIdentity{
|
||||
UserID: ¤tUser.ID,
|
||||
Provider: &provider,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get existing linked identity, error: %v", err)
|
||||
}
|
||||
if existingForProvider != nil {
|
||||
if existingForProvider.ExternUID == externUID {
|
||||
return currentUser, nil
|
||||
}
|
||||
return nil, status.Errorf(codes.AlreadyExists, "identity provider is already linked to another external account for this user")
|
||||
}
|
||||
|
||||
if _, err := s.Store.CreateUserIdentity(ctx, &store.UserIdentity{
|
||||
UserID: currentUser.ID,
|
||||
Provider: provider,
|
||||
ExternUID: externUID,
|
||||
}); err != nil {
|
||||
if isUniqueConstraintViolation(err) {
|
||||
winner, getErr := s.getLinkedSSOUser(ctx, provider, externUID)
|
||||
if getErr != nil {
|
||||
return nil, getErr
|
||||
}
|
||||
if winner != nil {
|
||||
if winner.ID != currentUser.ID {
|
||||
return nil, status.Errorf(codes.AlreadyExists, "identity provider account is already linked to another user")
|
||||
}
|
||||
return currentUser, nil
|
||||
}
|
||||
|
||||
existingForProvider, getErr := s.Store.GetUserIdentity(ctx, &store.FindUserIdentity{
|
||||
UserID: ¤tUser.ID,
|
||||
Provider: &provider,
|
||||
})
|
||||
if getErr != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to reload linked identity after race, error: %v", getErr)
|
||||
}
|
||||
if existingForProvider != nil {
|
||||
if existingForProvider.ExternUID == externUID {
|
||||
return currentUser, nil
|
||||
}
|
||||
return nil, status.Errorf(codes.AlreadyExists, "identity provider is already linked to another external account for this user")
|
||||
}
|
||||
|
||||
return nil, status.Errorf(codes.Internal, "user identity conflict reported but no winning row found")
|
||||
}
|
||||
return nil, status.Errorf(codes.Internal, "failed to create user identity, error: %v", err)
|
||||
}
|
||||
return currentUser, nil
|
||||
}
|
||||
|
||||
// isUniqueConstraintViolation matches the driver-specific error messages that each
|
||||
// supported backend emits when any UNIQUE constraint rejects an insert. Callers
|
||||
// disambiguate which constraint was hit from the insertion context (e.g. inserting
|
||||
// a user_identity row can only violate UNIQUE(provider, extern_uid); inserting a
|
||||
// user row can only violate UNIQUE(username)). Matches the pattern used in
|
||||
// memo_service.go for the memo UID unique check.
|
||||
func isUniqueConstraintViolation(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
msg := err.Error()
|
||||
return strings.Contains(msg, "UNIQUE constraint failed") ||
|
||||
strings.Contains(msg, "duplicate key") ||
|
||||
strings.Contains(msg, "Duplicate entry")
|
||||
}
|
||||
|
||||
// doSignIn performs the actual sign-in operation by creating a session and setting the cookie.
|
||||
//
|
||||
// This function:
|
||||
// 1. Generates refresh token and access token.
|
||||
// 2. Stores refresh token metadata in user_setting.
|
||||
// 3. Sets refresh token as HttpOnly cookie.
|
||||
// 4. Returns access token and its expiry time.
|
||||
func (s *APIV1Service) doSignIn(ctx context.Context, user *store.User) (string, time.Time, error) {
|
||||
// Generate refresh token
|
||||
tokenID := util.GenUUID()
|
||||
refreshToken, refreshExpiresAt, err := auth.GenerateRefreshToken(user.ID, tokenID, []byte(s.Secret))
|
||||
if err != nil {
|
||||
return "", time.Time{}, status.Errorf(codes.Internal, "failed to generate refresh token: %v", err)
|
||||
}
|
||||
|
||||
// Store refresh token metadata
|
||||
clientInfo := s.extractClientInfo(ctx)
|
||||
refreshTokenRecord := &storepb.RefreshTokensUserSetting_RefreshToken{
|
||||
TokenId: tokenID,
|
||||
ExpiresAt: timestamppb.New(refreshExpiresAt),
|
||||
CreatedAt: timestamppb.Now(),
|
||||
ClientInfo: clientInfo,
|
||||
}
|
||||
if err := s.Store.AddUserRefreshToken(ctx, user.ID, refreshTokenRecord); err != nil {
|
||||
slog.Error("failed to store refresh token", "error", err)
|
||||
}
|
||||
|
||||
// Set refresh token cookie
|
||||
refreshCookie := s.buildRefreshTokenCookie(ctx, refreshToken, refreshExpiresAt)
|
||||
if err := SetResponseHeader(ctx, "Set-Cookie", refreshCookie); err != nil {
|
||||
return "", time.Time{}, status.Errorf(codes.Internal, "failed to set refresh token cookie: %v", err)
|
||||
}
|
||||
|
||||
// Generate access token
|
||||
accessToken, accessExpiresAt, err := auth.GenerateAccessTokenV2(
|
||||
user.ID,
|
||||
user.Username,
|
||||
string(user.Role),
|
||||
string(user.RowStatus),
|
||||
[]byte(s.Secret),
|
||||
)
|
||||
if err != nil {
|
||||
return "", time.Time{}, status.Errorf(codes.Internal, "failed to generate access token: %v", err)
|
||||
}
|
||||
|
||||
return accessToken, accessExpiresAt, nil
|
||||
}
|
||||
|
||||
// SignOut terminates the user's authentication.
|
||||
// Revokes the refresh token and clears the authentication cookie.
|
||||
//
|
||||
// Authentication: Required (access token).
|
||||
// Returns: Empty response on success.
|
||||
func (s *APIV1Service) SignOut(ctx context.Context, _ *v1pb.SignOutRequest) (*emptypb.Empty, error) {
|
||||
// Get user from access token claims
|
||||
claims := auth.GetUserClaims(ctx)
|
||||
if claims != nil {
|
||||
// Revoke refresh token if we can identify it
|
||||
refreshToken := ""
|
||||
if md, ok := metadata.FromIncomingContext(ctx); ok {
|
||||
if cookies := md.Get("cookie"); len(cookies) > 0 {
|
||||
refreshToken = auth.ExtractRefreshTokenFromCookie(cookies[0])
|
||||
}
|
||||
}
|
||||
if refreshToken != "" {
|
||||
refreshClaims, err := auth.ParseRefreshToken(refreshToken, []byte(s.Secret))
|
||||
if err == nil {
|
||||
// Remove refresh token from user_setting by token_id
|
||||
_ = s.Store.RemoveUserRefreshToken(ctx, claims.UserID, refreshClaims.TokenID)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Clear refresh token cookie
|
||||
if err := s.clearAuthCookies(ctx); err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to clear auth cookies, error: %v", err)
|
||||
}
|
||||
return &emptypb.Empty{}, nil
|
||||
}
|
||||
|
||||
// RefreshToken exchanges a valid refresh token for a new access token.
|
||||
//
|
||||
// This endpoint implements refresh token rotation with sliding window sessions:
|
||||
// 1. Extracts the refresh token from the HttpOnly cookie (memos_refresh)
|
||||
// 2. Validates the refresh token against the database (checking expiry and revocation)
|
||||
// 3. Rotates the refresh token: generates a new one with fresh 30-day expiry
|
||||
// 4. Generates a new short-lived access token (15 minutes)
|
||||
// 5. Sets the new refresh token as HttpOnly cookie
|
||||
// 6. Returns the new access token and its expiry time
|
||||
//
|
||||
// Token rotation provides:
|
||||
// - Sliding window sessions: active users stay logged in indefinitely
|
||||
// - Better security: stolen refresh tokens become invalid after legitimate refresh
|
||||
//
|
||||
// Authentication: Requires valid refresh token in cookie (public endpoint)
|
||||
// Returns: New access token and expiry timestamp.
|
||||
func (s *APIV1Service) RefreshToken(ctx context.Context, _ *v1pb.RefreshTokenRequest) (*v1pb.RefreshTokenResponse, error) {
|
||||
// Extract refresh token from cookie
|
||||
refreshToken := ""
|
||||
if md, ok := metadata.FromIncomingContext(ctx); ok {
|
||||
if cookies := md.Get("cookie"); len(cookies) > 0 {
|
||||
refreshToken = auth.ExtractRefreshTokenFromCookie(cookies[0])
|
||||
}
|
||||
}
|
||||
|
||||
if refreshToken == "" {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "refresh token not found")
|
||||
}
|
||||
|
||||
// Validate refresh token and get old token ID for rotation
|
||||
authenticator := auth.NewAuthenticator(s.Store, s.Secret)
|
||||
user, oldTokenID, err := authenticator.AuthenticateByRefreshToken(ctx, refreshToken)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "invalid refresh token: %v", err)
|
||||
}
|
||||
|
||||
// --- Refresh Token Rotation ---
|
||||
// Generate new refresh token with fresh 30-day expiry (sliding window)
|
||||
newTokenID := util.GenUUID()
|
||||
newRefreshToken, newRefreshExpiresAt, err := auth.GenerateRefreshToken(user.ID, newTokenID, []byte(s.Secret))
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to generate refresh token: %v", err)
|
||||
}
|
||||
|
||||
// Store new refresh token (add before remove to handle race conditions)
|
||||
clientInfo := s.extractClientInfo(ctx)
|
||||
newRefreshTokenRecord := &storepb.RefreshTokensUserSetting_RefreshToken{
|
||||
TokenId: newTokenID,
|
||||
ExpiresAt: timestamppb.New(newRefreshExpiresAt),
|
||||
CreatedAt: timestamppb.Now(),
|
||||
ClientInfo: clientInfo,
|
||||
}
|
||||
if err := s.Store.AddUserRefreshToken(ctx, user.ID, newRefreshTokenRecord); err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to store refresh token: %v", err)
|
||||
}
|
||||
|
||||
// Remove old refresh token
|
||||
if err := s.Store.RemoveUserRefreshToken(ctx, user.ID, oldTokenID); err != nil {
|
||||
// Log but don't fail - old token will expire naturally
|
||||
slog.Warn("failed to remove old refresh token", "error", err, "userID", user.ID, "tokenID", oldTokenID)
|
||||
}
|
||||
|
||||
// Set new refresh token cookie
|
||||
newRefreshCookie := s.buildRefreshTokenCookie(ctx, newRefreshToken, newRefreshExpiresAt)
|
||||
if err := SetResponseHeader(ctx, "Set-Cookie", newRefreshCookie); err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to set refresh token cookie: %v", err)
|
||||
}
|
||||
// --- End Rotation ---
|
||||
|
||||
// Generate new access token
|
||||
accessToken, expiresAt, err := auth.GenerateAccessTokenV2(
|
||||
user.ID,
|
||||
user.Username,
|
||||
string(user.Role),
|
||||
string(user.RowStatus),
|
||||
[]byte(s.Secret),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to generate access token: %v", err)
|
||||
}
|
||||
|
||||
return &v1pb.RefreshTokenResponse{
|
||||
AccessToken: accessToken,
|
||||
ExpiresAt: timestamppb.New(expiresAt),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) clearAuthCookies(ctx context.Context) error {
|
||||
// Clear refresh token cookie
|
||||
refreshCookie := s.buildRefreshTokenCookie(ctx, "", time.Time{})
|
||||
if err := SetResponseHeader(ctx, "Set-Cookie", refreshCookie); err != nil {
|
||||
return errors.Wrap(err, "failed to set refresh cookie")
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func isSecureRequest(ctx context.Context) bool {
|
||||
md, ok := metadata.FromIncomingContext(ctx)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
|
||||
for _, value := range md.Get("x-forwarded-proto") {
|
||||
for _, proto := range strings.Split(value, ",") {
|
||||
if strings.EqualFold(strings.TrimSpace(proto), "https") {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for _, value := range md.Get("forwarded") {
|
||||
lowerValue := strings.ToLower(value)
|
||||
if strings.Contains(lowerValue, "proto=https") {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
for _, value := range md.Get("origin") {
|
||||
if strings.HasPrefix(strings.ToLower(strings.TrimSpace(value)), "https://") {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
func (*APIV1Service) buildRefreshTokenCookie(ctx context.Context, refreshToken string, expireTime time.Time) string {
|
||||
attrs := []string{
|
||||
fmt.Sprintf("%s=%s", auth.RefreshTokenCookieName, refreshToken),
|
||||
"Path=/",
|
||||
"HttpOnly",
|
||||
}
|
||||
if expireTime.IsZero() {
|
||||
attrs = append(attrs, "Expires=Thu, 01 Jan 1970 00:00:00 GMT")
|
||||
} else {
|
||||
// RFC 6265 requires cookie expiration dates to use GMT timezone
|
||||
// Convert to UTC and format with explicit "GMT" to ensure browser compatibility
|
||||
attrs = append(attrs, "Expires="+expireTime.UTC().Format("Mon, 02 Jan 2006 15:04:05 GMT"))
|
||||
}
|
||||
|
||||
if isSecureRequest(ctx) {
|
||||
attrs = append(attrs, "SameSite=Lax", "Secure")
|
||||
} else {
|
||||
attrs = append(attrs, "SameSite=Lax")
|
||||
}
|
||||
return strings.Join(attrs, "; ")
|
||||
}
|
||||
|
||||
func (s *APIV1Service) fetchCurrentUser(ctx context.Context) (*store.User, error) {
|
||||
userID := auth.GetUserID(ctx)
|
||||
if userID == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
user, err := s.Store.GetUser(ctx, &store.FindUser{
|
||||
ID: &userID,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if user == nil {
|
||||
return nil, errors.Errorf("user %d not found", userID)
|
||||
}
|
||||
if user.RowStatus == store.Archived {
|
||||
return nil, nil
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
|
||||
// extractClientInfo extracts comprehensive client information from the request context.
|
||||
//
|
||||
// This function parses metadata from the gRPC context to extract:
|
||||
// - User Agent: Raw user agent string for detailed parsing
|
||||
// - IP Address: Client IP from X-Forwarded-For or X-Real-IP headers
|
||||
// - Device Type: "mobile", "tablet", or "desktop" (parsed from user agent)
|
||||
// - Operating System: OS name and version (e.g., "iOS 17.1", "Windows 10/11")
|
||||
// - Browser: Browser name and version (e.g., "Chrome 120.0.0.0")
|
||||
//
|
||||
// This information enables users to:
|
||||
// - See all active sessions with device details
|
||||
// - Identify suspicious login attempts
|
||||
// - Revoke specific sessions from unknown devices.
|
||||
func (s *APIV1Service) extractClientInfo(ctx context.Context) *storepb.RefreshTokensUserSetting_ClientInfo {
|
||||
clientInfo := &storepb.RefreshTokensUserSetting_ClientInfo{}
|
||||
|
||||
// Extract user agent from metadata if available
|
||||
if md, ok := metadata.FromIncomingContext(ctx); ok {
|
||||
if userAgents := md.Get("user-agent"); len(userAgents) > 0 {
|
||||
userAgent := userAgents[0]
|
||||
clientInfo.UserAgent = userAgent
|
||||
|
||||
// Parse user agent to extract device type, OS, browser info
|
||||
s.parseUserAgent(userAgent, clientInfo)
|
||||
}
|
||||
if forwardedFor := md.Get("x-forwarded-for"); len(forwardedFor) > 0 {
|
||||
ipAddress := strings.Split(forwardedFor[0], ",")[0] // Get the first IP in case of multiple
|
||||
ipAddress = strings.TrimSpace(ipAddress)
|
||||
clientInfo.IpAddress = ipAddress
|
||||
} else if realIP := md.Get("x-real-ip"); len(realIP) > 0 {
|
||||
clientInfo.IpAddress = realIP[0]
|
||||
}
|
||||
}
|
||||
|
||||
return clientInfo
|
||||
}
|
||||
|
||||
// parseUserAgent extracts device type, OS, and browser information from user agent string.
|
||||
//
|
||||
// Detection logic:
|
||||
// - Device Type: Checks for keywords like "mobile", "tablet", "ipad"
|
||||
// - OS: Pattern matches for iOS, Android, Windows, macOS, Linux, Chrome OS
|
||||
// - Browser: Identifies Edge, Chrome, Firefox, Safari, Opera
|
||||
//
|
||||
// Note: This is a simplified parser. For production use with high accuracy requirements,
|
||||
// consider using a dedicated user agent parsing library.
|
||||
func (*APIV1Service) parseUserAgent(userAgent string, clientInfo *storepb.RefreshTokensUserSetting_ClientInfo) {
|
||||
if userAgent == "" {
|
||||
return
|
||||
}
|
||||
|
||||
userAgent = strings.ToLower(userAgent)
|
||||
|
||||
// Detect device type
|
||||
if strings.Contains(userAgent, "ipad") || strings.Contains(userAgent, "tablet") {
|
||||
clientInfo.DeviceType = "tablet"
|
||||
} else if strings.Contains(userAgent, "mobile") || strings.Contains(userAgent, "android") ||
|
||||
strings.Contains(userAgent, "iphone") || strings.Contains(userAgent, "ipod") ||
|
||||
strings.Contains(userAgent, "windows phone") || strings.Contains(userAgent, "blackberry") {
|
||||
clientInfo.DeviceType = "mobile"
|
||||
} else {
|
||||
clientInfo.DeviceType = "desktop"
|
||||
}
|
||||
|
||||
// Detect operating system
|
||||
if strings.Contains(userAgent, "iphone os") || strings.Contains(userAgent, "cpu os") {
|
||||
// Extract iOS version
|
||||
if idx := strings.Index(userAgent, "cpu os "); idx != -1 {
|
||||
versionStart := idx + 7
|
||||
versionEnd := strings.Index(userAgent[versionStart:], " ")
|
||||
if versionEnd != -1 {
|
||||
version := strings.ReplaceAll(userAgent[versionStart:versionStart+versionEnd], "_", ".")
|
||||
clientInfo.Os = "iOS " + version
|
||||
} else {
|
||||
clientInfo.Os = "iOS"
|
||||
}
|
||||
} else if idx := strings.Index(userAgent, "iphone os "); idx != -1 {
|
||||
versionStart := idx + 10
|
||||
versionEnd := strings.Index(userAgent[versionStart:], " ")
|
||||
if versionEnd != -1 {
|
||||
version := strings.ReplaceAll(userAgent[versionStart:versionStart+versionEnd], "_", ".")
|
||||
clientInfo.Os = "iOS " + version
|
||||
} else {
|
||||
clientInfo.Os = "iOS"
|
||||
}
|
||||
} else {
|
||||
clientInfo.Os = "iOS"
|
||||
}
|
||||
} else if strings.Contains(userAgent, "android") {
|
||||
// Extract Android version
|
||||
if idx := strings.Index(userAgent, "android "); idx != -1 {
|
||||
versionStart := idx + 8
|
||||
versionEnd := strings.Index(userAgent[versionStart:], ";")
|
||||
if versionEnd == -1 {
|
||||
versionEnd = strings.Index(userAgent[versionStart:], ")")
|
||||
}
|
||||
if versionEnd != -1 {
|
||||
version := userAgent[versionStart : versionStart+versionEnd]
|
||||
clientInfo.Os = "Android " + version
|
||||
} else {
|
||||
clientInfo.Os = "Android"
|
||||
}
|
||||
} else {
|
||||
clientInfo.Os = "Android"
|
||||
}
|
||||
} else if strings.Contains(userAgent, "windows nt 10.0") {
|
||||
clientInfo.Os = "Windows 10/11"
|
||||
} else if strings.Contains(userAgent, "windows nt 6.3") {
|
||||
clientInfo.Os = "Windows 8.1"
|
||||
} else if strings.Contains(userAgent, "windows nt 6.1") {
|
||||
clientInfo.Os = "Windows 7"
|
||||
} else if strings.Contains(userAgent, "windows") {
|
||||
clientInfo.Os = "Windows"
|
||||
} else if strings.Contains(userAgent, "mac os x") {
|
||||
// Extract macOS version
|
||||
if idx := strings.Index(userAgent, "mac os x "); idx != -1 {
|
||||
versionStart := idx + 9
|
||||
versionEnd := strings.Index(userAgent[versionStart:], ";")
|
||||
if versionEnd == -1 {
|
||||
versionEnd = strings.Index(userAgent[versionStart:], ")")
|
||||
}
|
||||
if versionEnd != -1 {
|
||||
version := strings.ReplaceAll(userAgent[versionStart:versionStart+versionEnd], "_", ".")
|
||||
clientInfo.Os = "macOS " + version
|
||||
} else {
|
||||
clientInfo.Os = "macOS"
|
||||
}
|
||||
} else {
|
||||
clientInfo.Os = "macOS"
|
||||
}
|
||||
} else if strings.Contains(userAgent, "linux") {
|
||||
clientInfo.Os = "Linux"
|
||||
} else if strings.Contains(userAgent, "cros") {
|
||||
clientInfo.Os = "Chrome OS"
|
||||
}
|
||||
|
||||
// Detect browser
|
||||
if strings.Contains(userAgent, "edg/") {
|
||||
// Extract Edge version
|
||||
if idx := strings.Index(userAgent, "edg/"); idx != -1 {
|
||||
versionStart := idx + 4
|
||||
versionEnd := strings.Index(userAgent[versionStart:], " ")
|
||||
if versionEnd == -1 {
|
||||
versionEnd = len(userAgent) - versionStart
|
||||
}
|
||||
version := userAgent[versionStart : versionStart+versionEnd]
|
||||
clientInfo.Browser = "Edge " + version
|
||||
} else {
|
||||
clientInfo.Browser = "Edge"
|
||||
}
|
||||
} else if strings.Contains(userAgent, "chrome/") && !strings.Contains(userAgent, "edg") {
|
||||
// Extract Chrome version
|
||||
if idx := strings.Index(userAgent, "chrome/"); idx != -1 {
|
||||
versionStart := idx + 7
|
||||
versionEnd := strings.Index(userAgent[versionStart:], " ")
|
||||
if versionEnd == -1 {
|
||||
versionEnd = len(userAgent) - versionStart
|
||||
}
|
||||
version := userAgent[versionStart : versionStart+versionEnd]
|
||||
clientInfo.Browser = "Chrome " + version
|
||||
} else {
|
||||
clientInfo.Browser = "Chrome"
|
||||
}
|
||||
} else if strings.Contains(userAgent, "firefox/") {
|
||||
// Extract Firefox version
|
||||
if idx := strings.Index(userAgent, "firefox/"); idx != -1 {
|
||||
versionStart := idx + 8
|
||||
versionEnd := strings.Index(userAgent[versionStart:], " ")
|
||||
if versionEnd == -1 {
|
||||
versionEnd = len(userAgent) - versionStart
|
||||
}
|
||||
version := userAgent[versionStart : versionStart+versionEnd]
|
||||
clientInfo.Browser = "Firefox " + version
|
||||
} else {
|
||||
clientInfo.Browser = "Firefox"
|
||||
}
|
||||
} else if strings.Contains(userAgent, "safari/") && !strings.Contains(userAgent, "chrome") && !strings.Contains(userAgent, "edg") {
|
||||
// Extract Safari version
|
||||
if idx := strings.Index(userAgent, "version/"); idx != -1 {
|
||||
versionStart := idx + 8
|
||||
versionEnd := strings.Index(userAgent[versionStart:], " ")
|
||||
if versionEnd == -1 {
|
||||
versionEnd = len(userAgent) - versionStart
|
||||
}
|
||||
version := userAgent[versionStart : versionStart+versionEnd]
|
||||
clientInfo.Browser = "Safari " + version
|
||||
} else {
|
||||
clientInfo.Browser = "Safari"
|
||||
}
|
||||
} else if strings.Contains(userAgent, "opera/") || strings.Contains(userAgent, "opr/") {
|
||||
clientInfo.Browser = "Opera"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,228 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"google.golang.org/grpc/metadata"
|
||||
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
)
|
||||
|
||||
func TestParseUserAgent(t *testing.T) {
|
||||
service := &APIV1Service{}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
userAgent string
|
||||
expectedDevice string
|
||||
expectedOS string
|
||||
expectedBrowser string
|
||||
}{
|
||||
{
|
||||
name: "Chrome on Windows",
|
||||
userAgent: "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/119.0.0.0 Safari/537.36",
|
||||
expectedDevice: "desktop",
|
||||
expectedOS: "Windows 10/11",
|
||||
expectedBrowser: "Chrome 119.0.0.0",
|
||||
},
|
||||
{
|
||||
name: "Safari on macOS",
|
||||
userAgent: "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/605.1.15 (KHTML, like Gecko) Version/17.0 Safari/605.1.15",
|
||||
expectedDevice: "desktop",
|
||||
expectedOS: "macOS 10.15.7",
|
||||
expectedBrowser: "Safari 17.0",
|
||||
},
|
||||
{
|
||||
name: "Chrome on Android Mobile",
|
||||
userAgent: "Mozilla/5.0 (Linux; Android 13; SM-G998B) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/119.0.0.0 Mobile Safari/537.36",
|
||||
expectedDevice: "mobile",
|
||||
expectedOS: "Android 13",
|
||||
expectedBrowser: "Chrome 119.0.0.0",
|
||||
},
|
||||
{
|
||||
name: "Safari on iPhone",
|
||||
userAgent: "Mozilla/5.0 (iPhone; CPU iPhone OS 17_0 like Mac OS X) AppleWebKit/605.1.15 (KHTML, like Gecko) Version/17.0 Mobile/15E148 Safari/604.1",
|
||||
expectedDevice: "mobile",
|
||||
expectedOS: "iOS 17.0",
|
||||
expectedBrowser: "Safari 17.0",
|
||||
},
|
||||
{
|
||||
name: "Firefox on Windows",
|
||||
userAgent: "Mozilla/5.0 (Windows NT 10.0; Win64; x64; rv:109.0) Gecko/20100101 Firefox/119.0",
|
||||
expectedDevice: "desktop",
|
||||
expectedOS: "Windows 10/11",
|
||||
expectedBrowser: "Firefox 119.0",
|
||||
},
|
||||
{
|
||||
name: "Edge on Windows",
|
||||
userAgent: "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/119.0.0.0 Safari/537.36 Edg/119.0.0.0",
|
||||
expectedDevice: "desktop",
|
||||
expectedOS: "Windows 10/11",
|
||||
expectedBrowser: "Edge 119.0.0.0",
|
||||
},
|
||||
{
|
||||
name: "iPad Safari",
|
||||
userAgent: "Mozilla/5.0 (iPad; CPU OS 17_0 like Mac OS X) AppleWebKit/605.1.15 (KHTML, like Gecko) Version/17.0 Mobile/15E148 Safari/604.1",
|
||||
expectedDevice: "tablet",
|
||||
expectedOS: "iOS 17.0",
|
||||
expectedBrowser: "Safari 17.0",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
clientInfo := &storepb.RefreshTokensUserSetting_ClientInfo{}
|
||||
service.parseUserAgent(tt.userAgent, clientInfo)
|
||||
|
||||
if clientInfo.DeviceType != tt.expectedDevice {
|
||||
t.Errorf("Expected device type %s, got %s", tt.expectedDevice, clientInfo.DeviceType)
|
||||
}
|
||||
if clientInfo.Os != tt.expectedOS {
|
||||
t.Errorf("Expected OS %s, got %s", tt.expectedOS, clientInfo.Os)
|
||||
}
|
||||
if clientInfo.Browser != tt.expectedBrowser {
|
||||
t.Errorf("Expected browser %s, got %s", tt.expectedBrowser, clientInfo.Browser)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractClientInfo(t *testing.T) {
|
||||
service := &APIV1Service{}
|
||||
|
||||
// Test with metadata containing user agent and IP
|
||||
md := metadata.New(map[string]string{
|
||||
"user-agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/119.0.0.0 Safari/537.36",
|
||||
"x-forwarded-for": "203.0.113.1, 198.51.100.1",
|
||||
"x-real-ip": "203.0.113.1",
|
||||
})
|
||||
|
||||
ctx := metadata.NewIncomingContext(context.Background(), md)
|
||||
|
||||
clientInfo := service.extractClientInfo(ctx)
|
||||
|
||||
if clientInfo.UserAgent == "" {
|
||||
t.Error("Expected user agent to be set")
|
||||
}
|
||||
if clientInfo.IpAddress != "203.0.113.1" {
|
||||
t.Errorf("Expected IP address to be 203.0.113.1, got %s", clientInfo.IpAddress)
|
||||
}
|
||||
if clientInfo.DeviceType != "desktop" {
|
||||
t.Errorf("Expected device type to be desktop, got %s", clientInfo.DeviceType)
|
||||
}
|
||||
if clientInfo.Os != "Windows 10/11" {
|
||||
t.Errorf("Expected OS to be Windows 10/11, got %s", clientInfo.Os)
|
||||
}
|
||||
if clientInfo.Browser != "Chrome 119.0.0.0" {
|
||||
t.Errorf("Expected browser to be Chrome 119.0.0.0, got %s", clientInfo.Browser)
|
||||
}
|
||||
}
|
||||
|
||||
// TestClientInfoExamples demonstrates the enhanced client info extraction with various user agents.
|
||||
func TestClientInfoExamples(t *testing.T) {
|
||||
service := &APIV1Service{}
|
||||
|
||||
examples := []struct {
|
||||
description string
|
||||
userAgent string
|
||||
}{
|
||||
{
|
||||
description: "Modern Chrome on Windows 11",
|
||||
userAgent: "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36",
|
||||
},
|
||||
{
|
||||
description: "Safari on iPhone 15 Pro",
|
||||
userAgent: "Mozilla/5.0 (iPhone; CPU iPhone OS 17_1 like Mac OS X) AppleWebKit/605.1.15 (KHTML, like Gecko) Version/17.1 Mobile/15E148 Safari/604.1",
|
||||
},
|
||||
{
|
||||
description: "Chrome on Samsung Galaxy",
|
||||
userAgent: "Mozilla/5.0 (Linux; Android 14; SM-S918B) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Mobile Safari/537.36",
|
||||
},
|
||||
{
|
||||
description: "Firefox on Ubuntu",
|
||||
userAgent: "Mozilla/5.0 (X11; Ubuntu; Linux x86_64; rv:109.0) Gecko/20100101 Firefox/120.0",
|
||||
},
|
||||
{
|
||||
description: "Edge on Windows 10",
|
||||
userAgent: "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36 Edg/120.0.0.0",
|
||||
},
|
||||
{
|
||||
description: "Safari on iPad Air",
|
||||
userAgent: "Mozilla/5.0 (iPad; CPU OS 17_1 like Mac OS X) AppleWebKit/605.1.15 (KHTML, like Gecko) Version/17.1 Mobile/15E148 Safari/604.1",
|
||||
},
|
||||
}
|
||||
|
||||
for _, example := range examples {
|
||||
t.Run(example.description, func(t *testing.T) {
|
||||
clientInfo := &storepb.RefreshTokensUserSetting_ClientInfo{}
|
||||
service.parseUserAgent(example.userAgent, clientInfo)
|
||||
|
||||
t.Logf("User Agent: %s", example.userAgent)
|
||||
t.Logf("Device Type: %s", clientInfo.DeviceType)
|
||||
t.Logf("Operating System: %s", clientInfo.Os)
|
||||
t.Logf("Browser: %s", clientInfo.Browser)
|
||||
t.Log("---")
|
||||
|
||||
// Ensure all fields are populated
|
||||
if clientInfo.DeviceType == "" {
|
||||
t.Error("Device type should not be empty")
|
||||
}
|
||||
if clientInfo.Os == "" {
|
||||
t.Error("OS should not be empty")
|
||||
}
|
||||
if clientInfo.Browser == "" {
|
||||
t.Error("Browser should not be empty")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildRefreshTokenCookieSecureFlag(t *testing.T) {
|
||||
service := &APIV1Service{}
|
||||
|
||||
t.Run("sets Secure for https origin", func(t *testing.T) {
|
||||
ctx := metadata.NewIncomingContext(context.Background(), metadata.Pairs(
|
||||
"origin", "https://memos.example",
|
||||
))
|
||||
cookie := service.buildRefreshTokenCookie(ctx, "token", testCookieExpiry())
|
||||
if !containsCookieAttribute(cookie, "Secure") {
|
||||
t.Fatalf("expected Secure attribute in cookie: %s", cookie)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("sets Secure for forwarded proto", func(t *testing.T) {
|
||||
ctx := metadata.NewIncomingContext(context.Background(), metadata.Pairs(
|
||||
"x-forwarded-proto", "https",
|
||||
))
|
||||
cookie := service.buildRefreshTokenCookie(ctx, "token", testCookieExpiry())
|
||||
if !containsCookieAttribute(cookie, "Secure") {
|
||||
t.Fatalf("expected Secure attribute in cookie: %s", cookie)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("omits Secure for plain http", func(t *testing.T) {
|
||||
ctx := metadata.NewIncomingContext(context.Background(), metadata.Pairs(
|
||||
"origin", "http://memos.example",
|
||||
))
|
||||
cookie := service.buildRefreshTokenCookie(ctx, "token", testCookieExpiry())
|
||||
if containsCookieAttribute(cookie, "Secure") {
|
||||
t.Fatalf("did not expect Secure attribute in cookie: %s", cookie)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func testCookieExpiry() time.Time {
|
||||
return time.Date(2030, time.January, 2, 3, 4, 5, 0, time.UTC)
|
||||
}
|
||||
|
||||
func containsCookieAttribute(cookie, attr string) bool {
|
||||
for _, part := range strings.Split(cookie, ";") {
|
||||
if strings.EqualFold(strings.TrimSpace(part), attr) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,83 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"encoding/base64"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
"google.golang.org/protobuf/proto"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
const (
|
||||
// DefaultPageSize is the default page size for requests.
|
||||
DefaultPageSize = 10
|
||||
// MaxPageSize is the maximum page size for requests.
|
||||
MaxPageSize = 1000
|
||||
)
|
||||
|
||||
func convertStateFromStore(rowStatus store.RowStatus) v1pb.State {
|
||||
switch rowStatus {
|
||||
case store.Normal:
|
||||
return v1pb.State_NORMAL
|
||||
case store.Archived:
|
||||
return v1pb.State_ARCHIVED
|
||||
default:
|
||||
return v1pb.State_STATE_UNSPECIFIED
|
||||
}
|
||||
}
|
||||
|
||||
func convertStateToStore(state v1pb.State) store.RowStatus {
|
||||
switch state {
|
||||
case v1pb.State_ARCHIVED:
|
||||
return store.Archived
|
||||
default:
|
||||
return store.Normal
|
||||
}
|
||||
}
|
||||
|
||||
func getPageToken(limit int, offset int) (string, error) {
|
||||
return marshalPageToken(&v1pb.PageToken{
|
||||
Limit: int32(limit),
|
||||
Offset: int32(offset),
|
||||
})
|
||||
}
|
||||
|
||||
func normalizePageSize(pageSize int32) int {
|
||||
limit := int(pageSize)
|
||||
if limit <= 0 {
|
||||
return DefaultPageSize
|
||||
}
|
||||
if limit > MaxPageSize {
|
||||
return MaxPageSize
|
||||
}
|
||||
return limit
|
||||
}
|
||||
|
||||
func marshalPageToken(pageToken *v1pb.PageToken) (string, error) {
|
||||
b, err := proto.Marshal(pageToken)
|
||||
if err != nil {
|
||||
return "", errors.Wrapf(err, "failed to marshal page token")
|
||||
}
|
||||
return base64.StdEncoding.EncodeToString(b), nil
|
||||
}
|
||||
|
||||
func unmarshalPageToken(s string, pageToken *v1pb.PageToken) error {
|
||||
b, err := base64.StdEncoding.DecodeString(s)
|
||||
if err != nil {
|
||||
return errors.Wrapf(err, "failed to decode page token")
|
||||
}
|
||||
if err := proto.Unmarshal(b, pageToken); err != nil {
|
||||
return errors.Wrapf(err, "failed to unmarshal page token")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func isSuperUser(user *store.User) bool {
|
||||
return user.Role == store.RoleAdmin
|
||||
}
|
||||
|
||||
func canModifyMemo(user *store.User, memo *store.Memo) bool {
|
||||
return user != nil && memo != nil && (memo.CreatorID == user.ID || isSuperUser(user))
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestNormalizePageSize(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
pageSize int32
|
||||
want int
|
||||
}{
|
||||
{
|
||||
name: "default for zero",
|
||||
pageSize: 0,
|
||||
want: DefaultPageSize,
|
||||
},
|
||||
{
|
||||
name: "default for negative",
|
||||
pageSize: -1,
|
||||
want: DefaultPageSize,
|
||||
},
|
||||
{
|
||||
name: "preserves valid size",
|
||||
pageSize: 42,
|
||||
want: 42,
|
||||
},
|
||||
{
|
||||
name: "clamps oversized size",
|
||||
pageSize: int32(MaxPageSize + 1),
|
||||
want: MaxPageSize,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
require.Equal(t, tt.want, normalizePageSize(tt.pageSize))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
"github.com/usememos/memos/proto/gen/api/v1/apiv1connect"
|
||||
)
|
||||
|
||||
// ConnectServiceHandler wraps APIV1Service to implement Connect handler interfaces.
|
||||
// It adapts the existing gRPC service implementations to work with Connect's
|
||||
// request/response wrapper types.
|
||||
//
|
||||
// This wrapper pattern allows us to:
|
||||
// - Reuse existing gRPC service implementations
|
||||
// - Support both native gRPC and Connect protocols
|
||||
// - Maintain a single source of truth for business logic.
|
||||
type ConnectServiceHandler struct {
|
||||
*APIV1Service
|
||||
}
|
||||
|
||||
// NewConnectServiceHandler creates a new Connect service handler.
|
||||
func NewConnectServiceHandler(svc *APIV1Service) *ConnectServiceHandler {
|
||||
return &ConnectServiceHandler{APIV1Service: svc}
|
||||
}
|
||||
|
||||
// RegisterConnectHandlers registers all Connect service handlers on the given mux.
|
||||
func (s *ConnectServiceHandler) RegisterConnectHandlers(mux *http.ServeMux, opts ...connect.HandlerOption) {
|
||||
// Register all service handlers
|
||||
handlers := []struct {
|
||||
path string
|
||||
handler http.Handler
|
||||
}{
|
||||
wrap(apiv1connect.NewInstanceServiceHandler(s, opts...)),
|
||||
wrap(apiv1connect.NewAuthServiceHandler(s, opts...)),
|
||||
wrap(apiv1connect.NewUserServiceHandler(s, opts...)),
|
||||
wrap(apiv1connect.NewMemoServiceHandler(s, opts...)),
|
||||
wrap(apiv1connect.NewAttachmentServiceHandler(s, opts...)),
|
||||
wrap(apiv1connect.NewAIServiceHandler(s, opts...)),
|
||||
wrap(apiv1connect.NewShortcutServiceHandler(s, opts...)),
|
||||
wrap(apiv1connect.NewIdentityProviderServiceHandler(s, opts...)),
|
||||
}
|
||||
|
||||
for _, h := range handlers {
|
||||
mux.Handle(h.path, h.handler)
|
||||
}
|
||||
}
|
||||
|
||||
// wrap converts (path, handler) return value to a struct for cleaner iteration.
|
||||
func wrap(path string, handler http.Handler) struct {
|
||||
path string
|
||||
handler http.Handler
|
||||
} {
|
||||
return struct {
|
||||
path string
|
||||
handler http.Handler
|
||||
}{path, handler}
|
||||
}
|
||||
|
||||
// convertGRPCError converts gRPC status errors to Connect errors.
|
||||
// This preserves the error code semantics between the two protocols.
|
||||
func convertGRPCError(err error) error {
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
if st, ok := status.FromError(err); ok {
|
||||
return connect.NewError(grpcCodeToConnectCode(st.Code()), err)
|
||||
}
|
||||
return connect.NewError(connect.CodeInternal, err)
|
||||
}
|
||||
|
||||
// grpcCodeToConnectCode converts gRPC status codes to Connect error codes.
|
||||
// gRPC and Connect use the same error code semantics, so this is a direct cast.
|
||||
// See: https://connectrpc.com/docs/protocol/#error-codes
|
||||
func grpcCodeToConnectCode(code codes.Code) connect.Code {
|
||||
return connect.Code(code)
|
||||
}
|
||||
@@ -0,0 +1,246 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"reflect"
|
||||
"runtime/debug"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
pkgerrors "github.com/pkg/errors"
|
||||
"google.golang.org/grpc/metadata"
|
||||
|
||||
"github.com/usememos/memos/server/auth"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
// MetadataInterceptor converts Connect HTTP headers to gRPC metadata.
|
||||
//
|
||||
// This ensures service methods can use metadata.FromIncomingContext() to access
|
||||
// headers like User-Agent, X-Forwarded-For, etc., regardless of whether the
|
||||
// request came via Connect RPC or gRPC-Gateway.
|
||||
type MetadataInterceptor struct{}
|
||||
|
||||
// NewMetadataInterceptor creates a new metadata interceptor.
|
||||
func NewMetadataInterceptor() *MetadataInterceptor {
|
||||
return &MetadataInterceptor{}
|
||||
}
|
||||
|
||||
func (*MetadataInterceptor) WrapUnary(next connect.UnaryFunc) connect.UnaryFunc {
|
||||
return func(ctx context.Context, req connect.AnyRequest) (connect.AnyResponse, error) {
|
||||
// Convert HTTP headers to gRPC metadata
|
||||
header := req.Header()
|
||||
md := metadata.MD{}
|
||||
|
||||
// Copy important headers for client info extraction
|
||||
if ua := header.Get("User-Agent"); ua != "" {
|
||||
md.Set("user-agent", ua)
|
||||
}
|
||||
if origin := header.Get("Origin"); origin != "" {
|
||||
md.Set("origin", origin)
|
||||
}
|
||||
if xff := header.Get("X-Forwarded-For"); xff != "" {
|
||||
md.Set("x-forwarded-for", xff)
|
||||
}
|
||||
if xfp := header.Get("X-Forwarded-Proto"); xfp != "" {
|
||||
md.Set("x-forwarded-proto", xfp)
|
||||
}
|
||||
if xri := header.Get("X-Real-Ip"); xri != "" {
|
||||
md.Set("x-real-ip", xri)
|
||||
}
|
||||
if forwarded := header.Get("Forwarded"); forwarded != "" {
|
||||
md.Set("forwarded", forwarded)
|
||||
}
|
||||
// Forward Cookie header for authentication methods that need it (e.g., RefreshToken)
|
||||
if cookie := header.Get("Cookie"); cookie != "" {
|
||||
md.Set("cookie", cookie)
|
||||
}
|
||||
|
||||
// Set metadata in context so services can use metadata.FromIncomingContext()
|
||||
ctx = metadata.NewIncomingContext(ctx, md)
|
||||
|
||||
// Execute the request
|
||||
resp, err := next(ctx, req)
|
||||
|
||||
// Prevent browser caching of API responses to avoid stale data issues
|
||||
// See: https://github.com/usememos/memos/issues/5470
|
||||
if !isNilAnyResponse(resp) && resp.Header() != nil {
|
||||
resp.Header().Set("Cache-Control", "no-cache, no-store, must-revalidate")
|
||||
resp.Header().Set("Pragma", "no-cache")
|
||||
resp.Header().Set("Expires", "0")
|
||||
}
|
||||
|
||||
return resp, err
|
||||
}
|
||||
}
|
||||
|
||||
func isNilAnyResponse(resp connect.AnyResponse) bool {
|
||||
if resp == nil {
|
||||
return true
|
||||
}
|
||||
val := reflect.ValueOf(resp)
|
||||
return val.Kind() == reflect.Ptr && val.IsNil()
|
||||
}
|
||||
|
||||
func (*MetadataInterceptor) WrapStreamingClient(next connect.StreamingClientFunc) connect.StreamingClientFunc {
|
||||
return next
|
||||
}
|
||||
|
||||
func (*MetadataInterceptor) WrapStreamingHandler(next connect.StreamingHandlerFunc) connect.StreamingHandlerFunc {
|
||||
return next
|
||||
}
|
||||
|
||||
// LoggingInterceptor logs Connect RPC requests with appropriate log levels.
|
||||
//
|
||||
// Log levels:
|
||||
// - INFO: Successful requests and expected client errors (not found, permission denied, etc.)
|
||||
// - ERROR: Server errors (internal, unavailable, etc.)
|
||||
type LoggingInterceptor struct {
|
||||
logStacktrace bool
|
||||
}
|
||||
|
||||
// NewLoggingInterceptor creates a new logging interceptor.
|
||||
func NewLoggingInterceptor(logStacktrace bool) *LoggingInterceptor {
|
||||
return &LoggingInterceptor{logStacktrace: logStacktrace}
|
||||
}
|
||||
|
||||
func (in *LoggingInterceptor) WrapUnary(next connect.UnaryFunc) connect.UnaryFunc {
|
||||
return func(ctx context.Context, req connect.AnyRequest) (connect.AnyResponse, error) {
|
||||
resp, err := next(ctx, req)
|
||||
in.log(req.Spec().Procedure, err)
|
||||
return resp, err
|
||||
}
|
||||
}
|
||||
|
||||
func (*LoggingInterceptor) WrapStreamingClient(next connect.StreamingClientFunc) connect.StreamingClientFunc {
|
||||
return next // No-op for server-side interceptor
|
||||
}
|
||||
|
||||
func (*LoggingInterceptor) WrapStreamingHandler(next connect.StreamingHandlerFunc) connect.StreamingHandlerFunc {
|
||||
return next // Streaming not used in this service
|
||||
}
|
||||
|
||||
func (in *LoggingInterceptor) log(procedure string, err error) {
|
||||
level, msg := in.classifyError(err)
|
||||
attrs := []slog.Attr{slog.String("method", procedure)}
|
||||
if err != nil {
|
||||
attrs = append(attrs, slog.String("error", err.Error()))
|
||||
if in.logStacktrace {
|
||||
attrs = append(attrs, slog.String("stacktrace", fmt.Sprintf("%+v", err)))
|
||||
}
|
||||
}
|
||||
slog.LogAttrs(context.Background(), level, msg, attrs...)
|
||||
}
|
||||
|
||||
func (*LoggingInterceptor) classifyError(err error) (slog.Level, string) {
|
||||
if err == nil {
|
||||
return slog.LevelInfo, "OK"
|
||||
}
|
||||
|
||||
var connectErr *connect.Error
|
||||
if !pkgerrors.As(err, &connectErr) {
|
||||
return slog.LevelError, "unknown error"
|
||||
}
|
||||
|
||||
// Client errors (expected, log at INFO)
|
||||
switch connectErr.Code() {
|
||||
case connect.CodeCanceled,
|
||||
connect.CodeInvalidArgument,
|
||||
connect.CodeNotFound,
|
||||
connect.CodeAlreadyExists,
|
||||
connect.CodePermissionDenied,
|
||||
connect.CodeUnauthenticated,
|
||||
connect.CodeResourceExhausted,
|
||||
connect.CodeFailedPrecondition,
|
||||
connect.CodeAborted,
|
||||
connect.CodeOutOfRange:
|
||||
return slog.LevelInfo, "client error"
|
||||
default:
|
||||
// Server errors
|
||||
return slog.LevelError, "server error"
|
||||
}
|
||||
}
|
||||
|
||||
// RecoveryInterceptor recovers from panics in Connect handlers and returns an internal error.
|
||||
type RecoveryInterceptor struct {
|
||||
logStacktrace bool
|
||||
}
|
||||
|
||||
// NewRecoveryInterceptor creates a new recovery interceptor.
|
||||
func NewRecoveryInterceptor(logStacktrace bool) *RecoveryInterceptor {
|
||||
return &RecoveryInterceptor{logStacktrace: logStacktrace}
|
||||
}
|
||||
|
||||
func (in *RecoveryInterceptor) WrapUnary(next connect.UnaryFunc) connect.UnaryFunc {
|
||||
return func(ctx context.Context, req connect.AnyRequest) (resp connect.AnyResponse, err error) {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
in.logPanic(req.Spec().Procedure, r)
|
||||
err = connect.NewError(connect.CodeInternal, pkgerrors.New("internal server error"))
|
||||
}
|
||||
}()
|
||||
return next(ctx, req)
|
||||
}
|
||||
}
|
||||
|
||||
func (*RecoveryInterceptor) WrapStreamingClient(next connect.StreamingClientFunc) connect.StreamingClientFunc {
|
||||
return next
|
||||
}
|
||||
|
||||
func (*RecoveryInterceptor) WrapStreamingHandler(next connect.StreamingHandlerFunc) connect.StreamingHandlerFunc {
|
||||
return next
|
||||
}
|
||||
|
||||
func (in *RecoveryInterceptor) logPanic(procedure string, panicValue any) {
|
||||
attrs := []slog.Attr{
|
||||
slog.String("method", procedure),
|
||||
slog.Any("panic", panicValue),
|
||||
}
|
||||
if in.logStacktrace {
|
||||
attrs = append(attrs, slog.String("stacktrace", string(debug.Stack())))
|
||||
}
|
||||
slog.LogAttrs(context.Background(), slog.LevelError, "panic recovered in Connect handler", attrs...)
|
||||
}
|
||||
|
||||
// AuthInterceptor handles authentication for Connect handlers.
|
||||
//
|
||||
// It enforces authentication for all endpoints except those listed in PublicMethods.
|
||||
// Role-based authorization (admin checks) remains in the service layer.
|
||||
type AuthInterceptor struct {
|
||||
authenticator *auth.Authenticator
|
||||
}
|
||||
|
||||
// NewAuthInterceptor creates a new auth interceptor.
|
||||
func NewAuthInterceptor(store *store.Store, secret string) *AuthInterceptor {
|
||||
return &AuthInterceptor{
|
||||
authenticator: auth.NewAuthenticator(store, secret),
|
||||
}
|
||||
}
|
||||
|
||||
func (in *AuthInterceptor) WrapUnary(next connect.UnaryFunc) connect.UnaryFunc {
|
||||
return func(ctx context.Context, req connect.AnyRequest) (connect.AnyResponse, error) {
|
||||
header := req.Header()
|
||||
authHeader := header.Get("Authorization")
|
||||
|
||||
result := in.authenticator.Authenticate(ctx, authHeader)
|
||||
|
||||
// Enforce authentication for non-public methods
|
||||
if result == nil && !IsPublicMethod(req.Spec().Procedure) {
|
||||
return nil, connect.NewError(connect.CodeUnauthenticated, errors.New("authentication required"))
|
||||
}
|
||||
|
||||
ctx = auth.ApplyToContext(ctx, result)
|
||||
|
||||
return next(ctx, req)
|
||||
}
|
||||
}
|
||||
|
||||
func (*AuthInterceptor) WrapStreamingClient(next connect.StreamingClientFunc) connect.StreamingClientFunc {
|
||||
return next
|
||||
}
|
||||
|
||||
func (*AuthInterceptor) WrapStreamingHandler(next connect.StreamingHandlerFunc) connect.StreamingHandlerFunc {
|
||||
return next
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
"google.golang.org/grpc/metadata"
|
||||
"google.golang.org/protobuf/types/known/emptypb"
|
||||
)
|
||||
|
||||
func TestMetadataInterceptorForwardsSecurityHeaders(t *testing.T) {
|
||||
interceptor := NewMetadataInterceptor()
|
||||
req := connect.NewRequest(&emptypb.Empty{})
|
||||
req.Header().Set("Origin", "https://memos.example")
|
||||
req.Header().Set("X-Forwarded-Proto", "https")
|
||||
req.Header().Set("Forwarded", "for=203.0.113.1;proto=https")
|
||||
|
||||
handler := interceptor.WrapUnary(func(ctx context.Context, _ connect.AnyRequest) (connect.AnyResponse, error) {
|
||||
md, ok := metadata.FromIncomingContext(ctx)
|
||||
if !ok {
|
||||
t.Fatal("expected metadata in context")
|
||||
}
|
||||
if got := md.Get("origin"); len(got) != 1 || got[0] != "https://memos.example" {
|
||||
t.Fatalf("unexpected origin metadata: %v", got)
|
||||
}
|
||||
if got := md.Get("x-forwarded-proto"); len(got) != 1 || got[0] != "https" {
|
||||
t.Fatalf("unexpected x-forwarded-proto metadata: %v", got)
|
||||
}
|
||||
if got := md.Get("forwarded"); len(got) != 1 || got[0] != "for=203.0.113.1;proto=https" {
|
||||
t.Fatalf("unexpected forwarded metadata: %v", got)
|
||||
}
|
||||
return connect.NewResponse(&emptypb.Empty{}), nil
|
||||
})
|
||||
|
||||
if _, err := handler(context.Background(), req); err != nil {
|
||||
t.Fatalf("metadata interceptor returned error: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,600 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
"google.golang.org/protobuf/types/known/emptypb"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
)
|
||||
|
||||
// This file contains all Connect service handler method implementations.
|
||||
// Each method delegates to the underlying gRPC service implementation,
|
||||
// converting between Connect and gRPC request/response types.
|
||||
|
||||
// InstanceService
|
||||
|
||||
func (s *ConnectServiceHandler) GetInstanceProfile(ctx context.Context, req *connect.Request[v1pb.GetInstanceProfileRequest]) (*connect.Response[v1pb.InstanceProfile], error) {
|
||||
resp, err := s.APIV1Service.GetInstanceProfile(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) GetInstanceSetting(ctx context.Context, req *connect.Request[v1pb.GetInstanceSettingRequest]) (*connect.Response[v1pb.InstanceSetting], error) {
|
||||
resp, err := s.APIV1Service.GetInstanceSetting(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) BatchGetInstanceSettings(ctx context.Context, req *connect.Request[v1pb.BatchGetInstanceSettingsRequest]) (*connect.Response[v1pb.BatchGetInstanceSettingsResponse], error) {
|
||||
resp, err := s.APIV1Service.BatchGetInstanceSettings(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) UpdateInstanceSetting(ctx context.Context, req *connect.Request[v1pb.UpdateInstanceSettingRequest]) (*connect.Response[v1pb.InstanceSetting], error) {
|
||||
resp, err := s.APIV1Service.UpdateInstanceSetting(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) TestInstanceEmailSetting(ctx context.Context, req *connect.Request[v1pb.TestInstanceEmailSettingRequest]) (*connect.Response[emptypb.Empty], error) {
|
||||
resp, err := s.APIV1Service.TestInstanceEmailSetting(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) GetInstanceStats(ctx context.Context, req *connect.Request[v1pb.GetInstanceStatsRequest]) (*connect.Response[v1pb.InstanceStats], error) {
|
||||
resp, err := s.APIV1Service.GetInstanceStats(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
// AuthService
|
||||
//
|
||||
// Auth service methods need special handling for response headers (cookies).
|
||||
// We use connectWithHeaderCarrier helper to inject a header carrier into the context,
|
||||
// which allows the service to set headers in a protocol-agnostic way.
|
||||
|
||||
func (s *ConnectServiceHandler) GetCurrentUser(ctx context.Context, req *connect.Request[v1pb.GetCurrentUserRequest]) (*connect.Response[v1pb.GetCurrentUserResponse], error) {
|
||||
return connectWithHeaderCarrier(ctx, func(ctx context.Context) (*v1pb.GetCurrentUserResponse, error) {
|
||||
return s.APIV1Service.GetCurrentUser(ctx, req.Msg)
|
||||
})
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) SignIn(ctx context.Context, req *connect.Request[v1pb.SignInRequest]) (*connect.Response[v1pb.SignInResponse], error) {
|
||||
return connectWithHeaderCarrier(ctx, func(ctx context.Context) (*v1pb.SignInResponse, error) {
|
||||
return s.APIV1Service.SignIn(ctx, req.Msg)
|
||||
})
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) SignOut(ctx context.Context, req *connect.Request[v1pb.SignOutRequest]) (*connect.Response[emptypb.Empty], error) {
|
||||
return connectWithHeaderCarrier(ctx, func(ctx context.Context) (*emptypb.Empty, error) {
|
||||
return s.APIV1Service.SignOut(ctx, req.Msg)
|
||||
})
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) RefreshToken(ctx context.Context, req *connect.Request[v1pb.RefreshTokenRequest]) (*connect.Response[v1pb.RefreshTokenResponse], error) {
|
||||
return connectWithHeaderCarrier(ctx, func(ctx context.Context) (*v1pb.RefreshTokenResponse, error) {
|
||||
return s.APIV1Service.RefreshToken(ctx, req.Msg)
|
||||
})
|
||||
}
|
||||
|
||||
// UserService
|
||||
|
||||
func (s *ConnectServiceHandler) ListUsers(ctx context.Context, req *connect.Request[v1pb.ListUsersRequest]) (*connect.Response[v1pb.ListUsersResponse], error) {
|
||||
resp, err := s.APIV1Service.ListUsers(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) BatchGetUsers(ctx context.Context, req *connect.Request[v1pb.BatchGetUsersRequest]) (*connect.Response[v1pb.BatchGetUsersResponse], error) {
|
||||
resp, err := s.APIV1Service.BatchGetUsers(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) GetUser(ctx context.Context, req *connect.Request[v1pb.GetUserRequest]) (*connect.Response[v1pb.User], error) {
|
||||
resp, err := s.APIV1Service.GetUser(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) CreateUser(ctx context.Context, req *connect.Request[v1pb.CreateUserRequest]) (*connect.Response[v1pb.User], error) {
|
||||
resp, err := s.APIV1Service.CreateUser(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) UpdateUser(ctx context.Context, req *connect.Request[v1pb.UpdateUserRequest]) (*connect.Response[v1pb.User], error) {
|
||||
resp, err := s.APIV1Service.UpdateUser(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) DeleteUser(ctx context.Context, req *connect.Request[v1pb.DeleteUserRequest]) (*connect.Response[emptypb.Empty], error) {
|
||||
return connectWithHeaderCarrier(ctx, func(ctx context.Context) (*emptypb.Empty, error) {
|
||||
return s.APIV1Service.DeleteUser(ctx, req.Msg)
|
||||
})
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) ListAllUserStats(ctx context.Context, req *connect.Request[v1pb.ListAllUserStatsRequest]) (*connect.Response[v1pb.ListAllUserStatsResponse], error) {
|
||||
resp, err := s.APIV1Service.ListAllUserStats(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) GetUserStats(ctx context.Context, req *connect.Request[v1pb.GetUserStatsRequest]) (*connect.Response[v1pb.UserStats], error) {
|
||||
resp, err := s.APIV1Service.GetUserStats(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) GetUserSetting(ctx context.Context, req *connect.Request[v1pb.GetUserSettingRequest]) (*connect.Response[v1pb.UserSetting], error) {
|
||||
resp, err := s.APIV1Service.GetUserSetting(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) UpdateUserSetting(ctx context.Context, req *connect.Request[v1pb.UpdateUserSettingRequest]) (*connect.Response[v1pb.UserSetting], error) {
|
||||
resp, err := s.APIV1Service.UpdateUserSetting(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) ListUserSettings(ctx context.Context, req *connect.Request[v1pb.ListUserSettingsRequest]) (*connect.Response[v1pb.ListUserSettingsResponse], error) {
|
||||
resp, err := s.APIV1Service.ListUserSettings(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) ListLinkedIdentities(ctx context.Context, req *connect.Request[v1pb.ListLinkedIdentitiesRequest]) (*connect.Response[v1pb.ListLinkedIdentitiesResponse], error) {
|
||||
resp, err := s.APIV1Service.ListLinkedIdentities(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) CreateLinkedIdentity(ctx context.Context, req *connect.Request[v1pb.CreateLinkedIdentityRequest]) (*connect.Response[v1pb.LinkedIdentity], error) {
|
||||
resp, err := s.APIV1Service.CreateLinkedIdentity(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) GetLinkedIdentity(ctx context.Context, req *connect.Request[v1pb.GetLinkedIdentityRequest]) (*connect.Response[v1pb.LinkedIdentity], error) {
|
||||
resp, err := s.APIV1Service.GetLinkedIdentity(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) DeleteLinkedIdentity(ctx context.Context, req *connect.Request[v1pb.DeleteLinkedIdentityRequest]) (*connect.Response[emptypb.Empty], error) {
|
||||
resp, err := s.APIV1Service.DeleteLinkedIdentity(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) ListPersonalAccessTokens(ctx context.Context, req *connect.Request[v1pb.ListPersonalAccessTokensRequest]) (*connect.Response[v1pb.ListPersonalAccessTokensResponse], error) {
|
||||
resp, err := s.APIV1Service.ListPersonalAccessTokens(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) CreatePersonalAccessToken(ctx context.Context, req *connect.Request[v1pb.CreatePersonalAccessTokenRequest]) (*connect.Response[v1pb.CreatePersonalAccessTokenResponse], error) {
|
||||
resp, err := s.APIV1Service.CreatePersonalAccessToken(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) DeletePersonalAccessToken(ctx context.Context, req *connect.Request[v1pb.DeletePersonalAccessTokenRequest]) (*connect.Response[emptypb.Empty], error) {
|
||||
resp, err := s.APIV1Service.DeletePersonalAccessToken(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) ListUserWebhooks(ctx context.Context, req *connect.Request[v1pb.ListUserWebhooksRequest]) (*connect.Response[v1pb.ListUserWebhooksResponse], error) {
|
||||
resp, err := s.APIV1Service.ListUserWebhooks(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) CreateUserWebhook(ctx context.Context, req *connect.Request[v1pb.CreateUserWebhookRequest]) (*connect.Response[v1pb.UserWebhook], error) {
|
||||
resp, err := s.APIV1Service.CreateUserWebhook(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) UpdateUserWebhook(ctx context.Context, req *connect.Request[v1pb.UpdateUserWebhookRequest]) (*connect.Response[v1pb.UserWebhook], error) {
|
||||
resp, err := s.APIV1Service.UpdateUserWebhook(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) DeleteUserWebhook(ctx context.Context, req *connect.Request[v1pb.DeleteUserWebhookRequest]) (*connect.Response[emptypb.Empty], error) {
|
||||
resp, err := s.APIV1Service.DeleteUserWebhook(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) ListUserNotifications(ctx context.Context, req *connect.Request[v1pb.ListUserNotificationsRequest]) (*connect.Response[v1pb.ListUserNotificationsResponse], error) {
|
||||
resp, err := s.APIV1Service.ListUserNotifications(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) UpdateUserNotification(ctx context.Context, req *connect.Request[v1pb.UpdateUserNotificationRequest]) (*connect.Response[v1pb.UserNotification], error) {
|
||||
resp, err := s.APIV1Service.UpdateUserNotification(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) DeleteUserNotification(ctx context.Context, req *connect.Request[v1pb.DeleteUserNotificationRequest]) (*connect.Response[emptypb.Empty], error) {
|
||||
resp, err := s.APIV1Service.DeleteUserNotification(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
// MemoService
|
||||
|
||||
func (s *ConnectServiceHandler) CreateMemo(ctx context.Context, req *connect.Request[v1pb.CreateMemoRequest]) (*connect.Response[v1pb.Memo], error) {
|
||||
resp, err := s.APIV1Service.CreateMemo(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) ListMemos(ctx context.Context, req *connect.Request[v1pb.ListMemosRequest]) (*connect.Response[v1pb.ListMemosResponse], error) {
|
||||
resp, err := s.APIV1Service.ListMemos(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) GetMemo(ctx context.Context, req *connect.Request[v1pb.GetMemoRequest]) (*connect.Response[v1pb.Memo], error) {
|
||||
resp, err := s.APIV1Service.GetMemo(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) UpdateMemo(ctx context.Context, req *connect.Request[v1pb.UpdateMemoRequest]) (*connect.Response[v1pb.Memo], error) {
|
||||
resp, err := s.APIV1Service.UpdateMemo(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) DeleteMemo(ctx context.Context, req *connect.Request[v1pb.DeleteMemoRequest]) (*connect.Response[emptypb.Empty], error) {
|
||||
resp, err := s.APIV1Service.DeleteMemo(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) SetMemoAttachments(ctx context.Context, req *connect.Request[v1pb.SetMemoAttachmentsRequest]) (*connect.Response[emptypb.Empty], error) {
|
||||
resp, err := s.APIV1Service.SetMemoAttachments(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) ListMemoAttachments(ctx context.Context, req *connect.Request[v1pb.ListMemoAttachmentsRequest]) (*connect.Response[v1pb.ListMemoAttachmentsResponse], error) {
|
||||
resp, err := s.APIV1Service.ListMemoAttachments(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) SetMemoRelations(ctx context.Context, req *connect.Request[v1pb.SetMemoRelationsRequest]) (*connect.Response[emptypb.Empty], error) {
|
||||
resp, err := s.APIV1Service.SetMemoRelations(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) ListMemoRelations(ctx context.Context, req *connect.Request[v1pb.ListMemoRelationsRequest]) (*connect.Response[v1pb.ListMemoRelationsResponse], error) {
|
||||
resp, err := s.APIV1Service.ListMemoRelations(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) CreateMemoComment(ctx context.Context, req *connect.Request[v1pb.CreateMemoCommentRequest]) (*connect.Response[v1pb.Memo], error) {
|
||||
resp, err := s.APIV1Service.CreateMemoComment(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) ListMemoComments(ctx context.Context, req *connect.Request[v1pb.ListMemoCommentsRequest]) (*connect.Response[v1pb.ListMemoCommentsResponse], error) {
|
||||
resp, err := s.APIV1Service.ListMemoComments(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) ListMemoReactions(ctx context.Context, req *connect.Request[v1pb.ListMemoReactionsRequest]) (*connect.Response[v1pb.ListMemoReactionsResponse], error) {
|
||||
resp, err := s.APIV1Service.ListMemoReactions(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) UpsertMemoReaction(ctx context.Context, req *connect.Request[v1pb.UpsertMemoReactionRequest]) (*connect.Response[v1pb.Reaction], error) {
|
||||
resp, err := s.APIV1Service.UpsertMemoReaction(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) DeleteMemoReaction(ctx context.Context, req *connect.Request[v1pb.DeleteMemoReactionRequest]) (*connect.Response[emptypb.Empty], error) {
|
||||
resp, err := s.APIV1Service.DeleteMemoReaction(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) CreateMemoShare(ctx context.Context, req *connect.Request[v1pb.CreateMemoShareRequest]) (*connect.Response[v1pb.MemoShare], error) {
|
||||
resp, err := s.APIV1Service.CreateMemoShare(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) ListMemoShares(ctx context.Context, req *connect.Request[v1pb.ListMemoSharesRequest]) (*connect.Response[v1pb.ListMemoSharesResponse], error) {
|
||||
resp, err := s.APIV1Service.ListMemoShares(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) DeleteMemoShare(ctx context.Context, req *connect.Request[v1pb.DeleteMemoShareRequest]) (*connect.Response[emptypb.Empty], error) {
|
||||
resp, err := s.APIV1Service.DeleteMemoShare(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) GetMemoByShare(ctx context.Context, req *connect.Request[v1pb.GetMemoByShareRequest]) (*connect.Response[v1pb.Memo], error) {
|
||||
resp, err := s.APIV1Service.GetMemoByShare(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) GetLinkMetadata(ctx context.Context, req *connect.Request[v1pb.GetLinkMetadataRequest]) (*connect.Response[v1pb.LinkMetadata], error) {
|
||||
resp, err := s.APIV1Service.GetLinkMetadata(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) BatchGetLinkMetadata(ctx context.Context, req *connect.Request[v1pb.BatchGetLinkMetadataRequest]) (*connect.Response[v1pb.BatchGetLinkMetadataResponse], error) {
|
||||
resp, err := s.APIV1Service.BatchGetLinkMetadata(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
// AttachmentService
|
||||
|
||||
func (s *ConnectServiceHandler) CreateAttachment(ctx context.Context, req *connect.Request[v1pb.CreateAttachmentRequest]) (*connect.Response[v1pb.Attachment], error) {
|
||||
resp, err := s.APIV1Service.CreateAttachment(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) ListAttachments(ctx context.Context, req *connect.Request[v1pb.ListAttachmentsRequest]) (*connect.Response[v1pb.ListAttachmentsResponse], error) {
|
||||
resp, err := s.APIV1Service.ListAttachments(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) GetAttachment(ctx context.Context, req *connect.Request[v1pb.GetAttachmentRequest]) (*connect.Response[v1pb.Attachment], error) {
|
||||
resp, err := s.APIV1Service.GetAttachment(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) UpdateAttachment(ctx context.Context, req *connect.Request[v1pb.UpdateAttachmentRequest]) (*connect.Response[v1pb.Attachment], error) {
|
||||
resp, err := s.APIV1Service.UpdateAttachment(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) DeleteAttachment(ctx context.Context, req *connect.Request[v1pb.DeleteAttachmentRequest]) (*connect.Response[emptypb.Empty], error) {
|
||||
resp, err := s.APIV1Service.DeleteAttachment(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) BatchDeleteAttachments(ctx context.Context, req *connect.Request[v1pb.BatchDeleteAttachmentsRequest]) (*connect.Response[emptypb.Empty], error) {
|
||||
resp, err := s.APIV1Service.BatchDeleteAttachments(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
// AIService
|
||||
|
||||
func (s *ConnectServiceHandler) Transcribe(ctx context.Context, req *connect.Request[v1pb.TranscribeRequest]) (*connect.Response[v1pb.TranscribeResponse], error) {
|
||||
resp, err := s.APIV1Service.Transcribe(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
// ShortcutService
|
||||
|
||||
func (s *ConnectServiceHandler) ListShortcuts(ctx context.Context, req *connect.Request[v1pb.ListShortcutsRequest]) (*connect.Response[v1pb.ListShortcutsResponse], error) {
|
||||
resp, err := s.APIV1Service.ListShortcuts(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) GetShortcut(ctx context.Context, req *connect.Request[v1pb.GetShortcutRequest]) (*connect.Response[v1pb.Shortcut], error) {
|
||||
resp, err := s.APIV1Service.GetShortcut(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) CreateShortcut(ctx context.Context, req *connect.Request[v1pb.CreateShortcutRequest]) (*connect.Response[v1pb.Shortcut], error) {
|
||||
resp, err := s.APIV1Service.CreateShortcut(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) UpdateShortcut(ctx context.Context, req *connect.Request[v1pb.UpdateShortcutRequest]) (*connect.Response[v1pb.Shortcut], error) {
|
||||
resp, err := s.APIV1Service.UpdateShortcut(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) DeleteShortcut(ctx context.Context, req *connect.Request[v1pb.DeleteShortcutRequest]) (*connect.Response[emptypb.Empty], error) {
|
||||
resp, err := s.APIV1Service.DeleteShortcut(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
// IdentityProviderService
|
||||
|
||||
func (s *ConnectServiceHandler) ListIdentityProviders(ctx context.Context, req *connect.Request[v1pb.ListIdentityProvidersRequest]) (*connect.Response[v1pb.ListIdentityProvidersResponse], error) {
|
||||
resp, err := s.APIV1Service.ListIdentityProviders(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) GetIdentityProvider(ctx context.Context, req *connect.Request[v1pb.GetIdentityProviderRequest]) (*connect.Response[v1pb.IdentityProvider], error) {
|
||||
resp, err := s.APIV1Service.GetIdentityProvider(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) CreateIdentityProvider(ctx context.Context, req *connect.Request[v1pb.CreateIdentityProviderRequest]) (*connect.Response[v1pb.IdentityProvider], error) {
|
||||
resp, err := s.APIV1Service.CreateIdentityProvider(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) UpdateIdentityProvider(ctx context.Context, req *connect.Request[v1pb.UpdateIdentityProviderRequest]) (*connect.Response[v1pb.IdentityProvider], error) {
|
||||
resp, err := s.APIV1Service.UpdateIdentityProvider(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
|
||||
func (s *ConnectServiceHandler) DeleteIdentityProvider(ctx context.Context, req *connect.Request[v1pb.DeleteIdentityProviderRequest]) (*connect.Response[emptypb.Empty], error) {
|
||||
resp, err := s.APIV1Service.DeleteIdentityProvider(ctx, req.Msg)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
return connect.NewResponse(resp), nil
|
||||
}
|
||||
@@ -0,0 +1,124 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/metadata"
|
||||
)
|
||||
|
||||
// headerCarrierKey is the context key for storing headers to be set in the response.
|
||||
type headerCarrierKey struct{}
|
||||
|
||||
// HeaderCarrier stores headers that need to be set in the response.
|
||||
//
|
||||
// Problem: The codebase supports two protocols simultaneously:
|
||||
// - Native gRPC: Uses grpc.SetHeader() to set response headers
|
||||
// - Connect-RPC: Uses connect.Response.Header().Set() to set response headers
|
||||
//
|
||||
// Solution: HeaderCarrier provides a protocol-agnostic way to set headers.
|
||||
// - Service methods call SetResponseHeader() regardless of protocol
|
||||
// - For gRPC requests: SetResponseHeader uses grpc.SetHeader directly
|
||||
// - For Connect requests: SetResponseHeader stores headers in HeaderCarrier
|
||||
// - Connect wrappers extract headers from HeaderCarrier and apply to response
|
||||
//
|
||||
// This allows service methods to work with both protocols without knowing which one is being used.
|
||||
type HeaderCarrier struct {
|
||||
headers map[string]string
|
||||
}
|
||||
|
||||
// newHeaderCarrier creates a new header carrier.
|
||||
func newHeaderCarrier() *HeaderCarrier {
|
||||
return &HeaderCarrier{
|
||||
headers: make(map[string]string),
|
||||
}
|
||||
}
|
||||
|
||||
// Set adds a header to the carrier.
|
||||
func (h *HeaderCarrier) Set(key, value string) {
|
||||
h.headers[key] = value
|
||||
}
|
||||
|
||||
// Get retrieves a header from the carrier.
|
||||
func (h *HeaderCarrier) Get(key string) string {
|
||||
return h.headers[key]
|
||||
}
|
||||
|
||||
// All returns all headers.
|
||||
func (h *HeaderCarrier) All() map[string]string {
|
||||
return h.headers
|
||||
}
|
||||
|
||||
// WithHeaderCarrier adds a header carrier to the context.
|
||||
func WithHeaderCarrier(ctx context.Context) context.Context {
|
||||
return context.WithValue(ctx, headerCarrierKey{}, newHeaderCarrier())
|
||||
}
|
||||
|
||||
// GetHeaderCarrier retrieves the header carrier from the context.
|
||||
// Returns nil if no carrier is present.
|
||||
func GetHeaderCarrier(ctx context.Context) *HeaderCarrier {
|
||||
if carrier, ok := ctx.Value(headerCarrierKey{}).(*HeaderCarrier); ok {
|
||||
return carrier
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetResponseHeader sets a header in the response.
|
||||
//
|
||||
// This function works for both gRPC and Connect protocols:
|
||||
// - For gRPC: Uses grpc.SetHeader to set headers in gRPC metadata
|
||||
// - For Connect: Stores in HeaderCarrier for Connect wrapper to apply later
|
||||
//
|
||||
// The protocol is automatically detected based on whether a HeaderCarrier
|
||||
// exists in the context (injected by Connect wrappers).
|
||||
func SetResponseHeader(ctx context.Context, key, value string) error {
|
||||
// Try Connect first (check if we have a header carrier)
|
||||
if carrier := GetHeaderCarrier(ctx); carrier != nil {
|
||||
carrier.Set(key, value)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Fall back to gRPC
|
||||
return grpc.SetHeader(ctx, metadata.New(map[string]string{
|
||||
key: value,
|
||||
}))
|
||||
}
|
||||
|
||||
// connectWithHeaderCarrier is a helper for Connect service wrappers that need to set response headers.
|
||||
//
|
||||
// It injects a HeaderCarrier into the context, calls the service method,
|
||||
// and applies any headers from the carrier to the Connect response.
|
||||
//
|
||||
// The generic parameter T is the non-pointer protobuf message type (e.g., v1pb.CreateSessionResponse),
|
||||
// while fn returns *T (the pointer type) as is standard for protobuf messages.
|
||||
//
|
||||
// Usage in Connect wrappers:
|
||||
//
|
||||
// func (s *ConnectServiceHandler) CreateSession(ctx context.Context, req *connect.Request[v1pb.CreateSessionRequest]) (*connect.Response[v1pb.CreateSessionResponse], error) {
|
||||
// return connectWithHeaderCarrier(ctx, func(ctx context.Context) (*v1pb.CreateSessionResponse, error) {
|
||||
// return s.APIV1Service.CreateSession(ctx, req.Msg)
|
||||
// })
|
||||
// }
|
||||
func connectWithHeaderCarrier[T any](ctx context.Context, fn func(context.Context) (*T, error)) (*connect.Response[T], error) {
|
||||
// Inject header carrier for Connect protocol
|
||||
ctx = WithHeaderCarrier(ctx)
|
||||
|
||||
// Call the service method
|
||||
resp, err := fn(ctx)
|
||||
if err != nil {
|
||||
return nil, convertGRPCError(err)
|
||||
}
|
||||
|
||||
// Create Connect response
|
||||
connectResp := connect.NewResponse(resp)
|
||||
|
||||
// Apply any headers set via the header carrier
|
||||
if carrier := GetHeaderCarrier(ctx); carrier != nil {
|
||||
for key, value := range carrier.All() {
|
||||
connectResp.Header().Set(key, value)
|
||||
}
|
||||
}
|
||||
|
||||
return connectResp, nil
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/health/grpc_health_v1"
|
||||
"google.golang.org/grpc/status"
|
||||
)
|
||||
|
||||
func (s *APIV1Service) Check(ctx context.Context,
|
||||
_ *grpc_health_v1.HealthCheckRequest) (*grpc_health_v1.HealthCheckResponse, error) {
|
||||
// Check if database is initialized by verifying instance basic setting exists
|
||||
instanceBasicSetting, err := s.Store.GetInstanceBasicSetting(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Unavailable, "database not initialized: %v", err)
|
||||
}
|
||||
|
||||
// Verify schema version is set (empty means database not properly initialized)
|
||||
if instanceBasicSetting.SchemaVersion == "" {
|
||||
return nil, status.Errorf(codes.Unavailable, "schema version not set")
|
||||
}
|
||||
|
||||
return &grpc_health_v1.HealthCheckResponse{Status: grpc_health_v1.HealthCheckResponse_SERVING}, nil
|
||||
}
|
||||
@@ -0,0 +1,235 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
"google.golang.org/protobuf/types/known/emptypb"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func (s *APIV1Service) CreateIdentityProvider(ctx context.Context, request *v1pb.CreateIdentityProviderRequest) (*v1pb.IdentityProvider, error) {
|
||||
currentUser, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get user: %v", err)
|
||||
}
|
||||
if currentUser == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
if currentUser.Role != store.RoleAdmin {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
|
||||
idpUID, err := ValidateAndGenerateUID(request.IdentityProviderId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
storeIdp := convertIdentityProviderToStore(request.IdentityProvider)
|
||||
storeIdp.Uid = idpUID
|
||||
|
||||
identityProvider, err := s.Store.CreateIdentityProvider(ctx, storeIdp)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to create identity provider, error: %+v", err)
|
||||
}
|
||||
return convertIdentityProviderFromStore(identityProvider), nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) ListIdentityProviders(ctx context.Context, _ *v1pb.ListIdentityProvidersRequest) (*v1pb.ListIdentityProvidersResponse, error) {
|
||||
identityProviders, err := s.Store.ListIdentityProviders(ctx, &store.FindIdentityProvider{})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to list identity providers, error: %+v", err)
|
||||
}
|
||||
|
||||
response := &v1pb.ListIdentityProvidersResponse{
|
||||
IdentityProviders: []*v1pb.IdentityProvider{},
|
||||
}
|
||||
for _, identityProvider := range identityProviders {
|
||||
response.IdentityProviders = append(response.IdentityProviders, convertIdentityProviderFromStore(identityProvider))
|
||||
}
|
||||
return response, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) GetIdentityProvider(ctx context.Context, request *v1pb.GetIdentityProviderRequest) (*v1pb.IdentityProvider, error) {
|
||||
uid, err := ExtractIdentityProviderUIDFromName(request.Name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid identity provider name: %v", err)
|
||||
}
|
||||
identityProvider, err := s.Store.GetIdentityProvider(ctx, &store.FindIdentityProvider{
|
||||
UID: &uid,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get identity provider, error: %+v", err)
|
||||
}
|
||||
if identityProvider == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "identity provider not found")
|
||||
}
|
||||
|
||||
return convertIdentityProviderFromStore(identityProvider), nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) UpdateIdentityProvider(ctx context.Context, request *v1pb.UpdateIdentityProviderRequest) (*v1pb.IdentityProvider, error) {
|
||||
currentUser, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get user: %v", err)
|
||||
}
|
||||
if currentUser == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
if currentUser.Role != store.RoleAdmin {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
|
||||
if request.UpdateMask == nil || len(request.UpdateMask.Paths) == 0 {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "update_mask is required")
|
||||
}
|
||||
|
||||
uid, err := ExtractIdentityProviderUIDFromName(request.IdentityProvider.Name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid identity provider name: %v", err)
|
||||
}
|
||||
|
||||
// Look up the IdP by UID to get the internal ID for update.
|
||||
existing, err := s.Store.GetIdentityProvider(ctx, &store.FindIdentityProvider{UID: &uid})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get identity provider, error: %+v", err)
|
||||
}
|
||||
if existing == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "identity provider not found")
|
||||
}
|
||||
|
||||
update := &store.UpdateIdentityProviderV1{
|
||||
ID: existing.Id,
|
||||
Type: storepb.IdentityProvider_Type(storepb.IdentityProvider_Type_value[request.IdentityProvider.Type.String()]),
|
||||
}
|
||||
for _, field := range request.UpdateMask.Paths {
|
||||
switch field {
|
||||
case "title":
|
||||
update.Name = &request.IdentityProvider.Title
|
||||
case "identifier_filter":
|
||||
update.IdentifierFilter = &request.IdentityProvider.IdentifierFilter
|
||||
case "config":
|
||||
update.Config = convertIdentityProviderConfigToStore(request.IdentityProvider.Type, request.IdentityProvider.Config)
|
||||
default:
|
||||
// Ignore unsupported fields
|
||||
}
|
||||
}
|
||||
|
||||
// Preserve write-only credential when the caller sends an empty value.
|
||||
if update.Config != nil {
|
||||
if oauth2Config := update.Config.GetOauth2Config(); oauth2Config != nil && oauth2Config.ClientSecret == "" {
|
||||
if existingOAuth := existing.Config.GetOauth2Config(); existingOAuth != nil {
|
||||
oauth2Config.ClientSecret = existingOAuth.ClientSecret
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
identityProvider, err := s.Store.UpdateIdentityProvider(ctx, update)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to update identity provider, error: %+v", err)
|
||||
}
|
||||
return convertIdentityProviderFromStore(identityProvider), nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) DeleteIdentityProvider(ctx context.Context, request *v1pb.DeleteIdentityProviderRequest) (*emptypb.Empty, error) {
|
||||
currentUser, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get user: %v", err)
|
||||
}
|
||||
if currentUser == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
if currentUser.Role != store.RoleAdmin {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
|
||||
uid, err := ExtractIdentityProviderUIDFromName(request.Name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid identity provider name: %v", err)
|
||||
}
|
||||
|
||||
// Look up the IdP by UID to get the internal ID for deletion.
|
||||
identityProvider, err := s.Store.GetIdentityProvider(ctx, &store.FindIdentityProvider{UID: &uid})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to check identity provider existence: %v", err)
|
||||
}
|
||||
if identityProvider == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "identity provider not found")
|
||||
}
|
||||
|
||||
if err := s.Store.DeleteIdentityProvider(ctx, &store.DeleteIdentityProvider{ID: identityProvider.Id}); err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to delete identity provider, error: %+v", err)
|
||||
}
|
||||
return &emptypb.Empty{}, nil
|
||||
}
|
||||
|
||||
func convertIdentityProviderFromStore(identityProvider *storepb.IdentityProvider) *v1pb.IdentityProvider {
|
||||
temp := &v1pb.IdentityProvider{
|
||||
Name: fmt.Sprintf("%s%s", IdentityProviderNamePrefix, identityProvider.Uid),
|
||||
Title: identityProvider.Name,
|
||||
IdentifierFilter: identityProvider.IdentifierFilter,
|
||||
Type: v1pb.IdentityProvider_Type(v1pb.IdentityProvider_Type_value[identityProvider.Type.String()]),
|
||||
}
|
||||
if identityProvider.Type == storepb.IdentityProvider_OAUTH2 {
|
||||
oauth2Config := identityProvider.Config.GetOauth2Config()
|
||||
temp.Config = &v1pb.IdentityProviderConfig{
|
||||
Config: &v1pb.IdentityProviderConfig_Oauth2Config{
|
||||
Oauth2Config: &v1pb.OAuth2Config{
|
||||
ClientId: oauth2Config.ClientId,
|
||||
// ClientSecret is write-only: never returned in responses.
|
||||
AuthUrl: oauth2Config.AuthUrl,
|
||||
TokenUrl: oauth2Config.TokenUrl,
|
||||
UserInfoUrl: oauth2Config.UserInfoUrl,
|
||||
Scopes: oauth2Config.Scopes,
|
||||
FieldMapping: &v1pb.FieldMapping{
|
||||
Identifier: oauth2Config.FieldMapping.Identifier,
|
||||
DisplayName: oauth2Config.FieldMapping.DisplayName,
|
||||
Email: oauth2Config.FieldMapping.Email,
|
||||
AvatarUrl: oauth2Config.FieldMapping.AvatarUrl,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
return temp
|
||||
}
|
||||
|
||||
func convertIdentityProviderToStore(identityProvider *v1pb.IdentityProvider) *storepb.IdentityProvider {
|
||||
temp := &storepb.IdentityProvider{
|
||||
Name: identityProvider.Title,
|
||||
IdentifierFilter: identityProvider.IdentifierFilter,
|
||||
Type: storepb.IdentityProvider_Type(storepb.IdentityProvider_Type_value[identityProvider.Type.String()]),
|
||||
Config: convertIdentityProviderConfigToStore(identityProvider.Type, identityProvider.Config),
|
||||
}
|
||||
return temp
|
||||
}
|
||||
|
||||
func convertIdentityProviderConfigToStore(identityProviderType v1pb.IdentityProvider_Type, config *v1pb.IdentityProviderConfig) *storepb.IdentityProviderConfig {
|
||||
if identityProviderType == v1pb.IdentityProvider_OAUTH2 {
|
||||
oauth2Config := config.GetOauth2Config()
|
||||
return &storepb.IdentityProviderConfig{
|
||||
Config: &storepb.IdentityProviderConfig_Oauth2Config{
|
||||
Oauth2Config: &storepb.OAuth2Config{
|
||||
ClientId: oauth2Config.ClientId,
|
||||
ClientSecret: oauth2Config.ClientSecret,
|
||||
AuthUrl: oauth2Config.AuthUrl,
|
||||
TokenUrl: oauth2Config.TokenUrl,
|
||||
UserInfoUrl: oauth2Config.UserInfoUrl,
|
||||
Scopes: oauth2Config.Scopes,
|
||||
FieldMapping: &storepb.FieldMapping{
|
||||
Identifier: oauth2Config.FieldMapping.Identifier,
|
||||
DisplayName: oauth2Config.FieldMapping.DisplayName,
|
||||
Email: oauth2Config.FieldMapping.Email,
|
||||
AvatarUrl: oauth2Config.FieldMapping.AvatarUrl,
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,836 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"math"
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
"github.com/lithammer/shortuuid/v4"
|
||||
"github.com/pkg/errors"
|
||||
colorpb "google.golang.org/genproto/googleapis/type/color"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
"google.golang.org/protobuf/types/known/emptypb"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/server/notification"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
const (
|
||||
maxTranscriptionConfigModelLength = 256
|
||||
maxTranscriptionConfigLanguageLength = 32
|
||||
maxTranscriptionConfigPromptLength = 4096
|
||||
maxBatchGetInstanceSettings = 100
|
||||
)
|
||||
|
||||
type instanceSettingCaller struct {
|
||||
user *store.User
|
||||
loaded bool
|
||||
}
|
||||
|
||||
func (c *instanceSettingCaller) currentUser(ctx context.Context, service *APIV1Service) (*store.User, error) {
|
||||
if c.loaded {
|
||||
return c.user, nil
|
||||
}
|
||||
user, err := service.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
c.user = user
|
||||
c.loaded = true
|
||||
return c.user, nil
|
||||
}
|
||||
|
||||
// GetInstanceProfile returns the instance profile.
|
||||
func (s *APIV1Service) GetInstanceProfile(ctx context.Context, _ *v1pb.GetInstanceProfileRequest) (*v1pb.InstanceProfile, error) {
|
||||
admin, err := s.GetInstanceAdmin(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get instance admin: %v", err)
|
||||
}
|
||||
|
||||
instanceProfile := &v1pb.InstanceProfile{
|
||||
Version: s.Profile.Version,
|
||||
Demo: s.Profile.Demo,
|
||||
InstanceUrl: s.Profile.InstanceURL,
|
||||
Admin: admin, // nil when not initialized
|
||||
Commit: s.Profile.Commit,
|
||||
}
|
||||
return instanceProfile, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) GetInstanceSetting(ctx context.Context, request *v1pb.GetInstanceSettingRequest) (*v1pb.InstanceSetting, error) {
|
||||
return s.getInstanceSettingByName(ctx, request.Name, &instanceSettingCaller{})
|
||||
}
|
||||
|
||||
// BatchGetInstanceSettings returns multiple instance settings in request order.
|
||||
func (s *APIV1Service) BatchGetInstanceSettings(ctx context.Context, request *v1pb.BatchGetInstanceSettingsRequest) (*v1pb.BatchGetInstanceSettingsResponse, error) {
|
||||
if len(request.Names) > maxBatchGetInstanceSettings {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "too many instance setting names (max %d)", maxBatchGetInstanceSettings)
|
||||
}
|
||||
|
||||
caller := &instanceSettingCaller{}
|
||||
settings := make([]*v1pb.InstanceSetting, 0, len(request.Names))
|
||||
for _, name := range request.Names {
|
||||
setting, err := s.getInstanceSettingByName(ctx, name, caller)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
settings = append(settings, setting)
|
||||
}
|
||||
|
||||
return &v1pb.BatchGetInstanceSettingsResponse{Settings: settings}, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) getInstanceSettingByName(ctx context.Context, name string, caller *instanceSettingCaller) (*v1pb.InstanceSetting, error) {
|
||||
instanceSettingKeyString, err := ExtractInstanceSettingKeyFromName(name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid instance setting name: %v", err)
|
||||
}
|
||||
|
||||
instanceSettingKey := storepb.InstanceSettingKey(storepb.InstanceSettingKey_value[instanceSettingKeyString])
|
||||
// Get instance setting from store with default value.
|
||||
switch instanceSettingKey {
|
||||
case storepb.InstanceSettingKey_BASIC:
|
||||
_, err = s.Store.GetInstanceBasicSetting(ctx)
|
||||
case storepb.InstanceSettingKey_GENERAL:
|
||||
_, err = s.Store.GetInstanceGeneralSetting(ctx)
|
||||
case storepb.InstanceSettingKey_MEMO_RELATED:
|
||||
_, err = s.Store.GetInstanceMemoRelatedSetting(ctx)
|
||||
case storepb.InstanceSettingKey_STORAGE:
|
||||
_, err = s.Store.GetInstanceStorageSetting(ctx)
|
||||
case storepb.InstanceSettingKey_TAGS:
|
||||
_, err = s.Store.GetInstanceTagsSetting(ctx)
|
||||
case storepb.InstanceSettingKey_NOTIFICATION:
|
||||
_, err = s.Store.GetInstanceNotificationSetting(ctx)
|
||||
case storepb.InstanceSettingKey_AI:
|
||||
_, err = s.Store.GetInstanceAISetting(ctx)
|
||||
default:
|
||||
return nil, status.Errorf(codes.InvalidArgument, "unsupported instance setting key: %v", instanceSettingKey)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get instance setting: %v", err)
|
||||
}
|
||||
|
||||
instanceSetting, err := s.Store.GetInstanceSetting(ctx, &store.FindInstanceSetting{
|
||||
Name: instanceSettingKey.String(),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get instance setting: %v", err)
|
||||
}
|
||||
if instanceSetting == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "instance setting not found")
|
||||
}
|
||||
|
||||
// Storage and notification settings contain credentials; restrict to admins only.
|
||||
if instanceSetting.Key == storepb.InstanceSettingKey_STORAGE ||
|
||||
instanceSetting.Key == storepb.InstanceSettingKey_NOTIFICATION {
|
||||
user, err := caller.currentUser(ctx, s)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
if user.Role != store.RoleAdmin {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
}
|
||||
isAdminCaller := false
|
||||
if instanceSetting.Key == storepb.InstanceSettingKey_AI {
|
||||
user, err := caller.currentUser(ctx, s)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
isAdminCaller = user.Role == store.RoleAdmin
|
||||
}
|
||||
|
||||
result := convertInstanceSettingFromStore(instanceSetting)
|
||||
if instanceSetting.Key == storepb.InstanceSettingKey_AI && !isAdminCaller {
|
||||
// Non-admin callers only need transcription.provider_id to gate the
|
||||
// editor's Transcribe button. Model / language / prompt are
|
||||
// admin-entered defaults that may contain proprietary glossary terms,
|
||||
// so they are redacted from non-admin responses.
|
||||
if ai := result.GetAiSetting(); ai != nil && ai.Transcription != nil {
|
||||
ai.Transcription.Model = ""
|
||||
ai.Transcription.Language = ""
|
||||
ai.Transcription.Prompt = ""
|
||||
}
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) UpdateInstanceSetting(ctx context.Context, request *v1pb.UpdateInstanceSettingRequest) (*v1pb.InstanceSetting, error) {
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
if user.Role != store.RoleAdmin {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
|
||||
// TODO: Apply update_mask if specified
|
||||
_ = request.UpdateMask
|
||||
|
||||
if err := validateInstanceSetting(request.Setting); err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid instance setting: %v", err)
|
||||
}
|
||||
|
||||
updateSetting := convertInstanceSettingToStore(request.Setting)
|
||||
|
||||
// Preserve write-only credential fields when the caller sends an empty value.
|
||||
// An empty string means "no change", not "clear the credential".
|
||||
switch updateSetting.Key {
|
||||
case storepb.InstanceSettingKey_NOTIFICATION:
|
||||
if notif := updateSetting.GetNotificationSetting(); notif != nil && notif.Email != nil && notif.Email.SmtpPassword == "" {
|
||||
existing, err := s.Store.GetInstanceNotificationSetting(ctx)
|
||||
if err == nil && existing != nil && existing.Email != nil {
|
||||
if existing.Email.SmtpPassword != "" && !sameSMTPConnectionIdentity(notif.Email, existing.Email) {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "smtp password is required when changing SMTP host, port, username, or encryption settings")
|
||||
}
|
||||
notif.Email.SmtpPassword = existing.Email.SmtpPassword
|
||||
}
|
||||
}
|
||||
case storepb.InstanceSettingKey_STORAGE:
|
||||
if storage := updateSetting.GetStorageSetting(); storage != nil && storage.S3Config != nil && storage.S3Config.AccessKeySecret == "" {
|
||||
existing, err := s.Store.GetInstanceStorageSetting(ctx)
|
||||
if err == nil && existing != nil && existing.S3Config != nil {
|
||||
storage.S3Config.AccessKeySecret = existing.S3Config.AccessKeySecret
|
||||
}
|
||||
}
|
||||
case storepb.InstanceSettingKey_AI:
|
||||
if err := s.prepareInstanceAISettingForUpdate(ctx, updateSetting.GetAiSetting()); err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid AI setting: %v", err)
|
||||
}
|
||||
default:
|
||||
// No credential preservation needed for other setting types.
|
||||
}
|
||||
|
||||
instanceSetting, err := s.Store.UpsertInstanceSetting(ctx, updateSetting)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to upsert instance setting: %v", err)
|
||||
}
|
||||
|
||||
return convertInstanceSettingFromStore(instanceSetting), nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) TestInstanceEmailSetting(ctx context.Context, request *v1pb.TestInstanceEmailSettingRequest) (*emptypb.Empty, error) {
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
if user.Role != store.RoleAdmin {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
|
||||
emailSetting, err := s.resolveTestEmailSetting(ctx, request.Email)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
recipientEmail := strings.TrimSpace(request.RecipientEmail)
|
||||
if recipientEmail == "" {
|
||||
recipientEmail = strings.TrimSpace(user.Email)
|
||||
}
|
||||
if recipientEmail == "" {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "recipient email is required")
|
||||
}
|
||||
|
||||
if err := notification.ValidateEmailSetting(emailSetting); err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid notification email setting: %v", err)
|
||||
}
|
||||
|
||||
if err := notification.SendTestEmail(emailSetting, recipientEmail); err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to send test email: %v. Check that the SMTP port matches encryption: Gmail uses port 587 with STARTTLS on and SSL/TLS off; port 465 requires SSL/TLS on", err)
|
||||
}
|
||||
|
||||
return &emptypb.Empty{}, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) resolveTestEmailSetting(ctx context.Context, requestEmail *v1pb.InstanceSetting_NotificationSetting_EmailSetting) (*storepb.InstanceNotificationSetting_EmailSetting, error) {
|
||||
if requestEmail == nil {
|
||||
existing, err := s.Store.GetInstanceNotificationSetting(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get notification setting: %v", err)
|
||||
}
|
||||
return existing.GetEmail(), nil
|
||||
}
|
||||
|
||||
emailSetting := convertInstanceNotificationSettingToStore(&v1pb.InstanceSetting_NotificationSetting{Email: requestEmail}).GetEmail()
|
||||
if emailSetting.SmtpPassword != "" {
|
||||
return emailSetting, nil
|
||||
}
|
||||
|
||||
existing, err := s.Store.GetInstanceNotificationSetting(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get notification setting: %v", err)
|
||||
}
|
||||
existingEmail := existing.GetEmail()
|
||||
if existingEmail == nil || existingEmail.SmtpPassword == "" {
|
||||
return emailSetting, nil
|
||||
}
|
||||
if sameSMTPConnectionIdentity(emailSetting, existingEmail) {
|
||||
emailSetting.SmtpPassword = existingEmail.SmtpPassword
|
||||
return emailSetting, nil
|
||||
}
|
||||
return nil, status.Errorf(codes.InvalidArgument, "smtp password is required when changing SMTP host, port, username, or encryption settings")
|
||||
}
|
||||
|
||||
func sameSMTPConnectionIdentity(setting, existing *storepb.InstanceNotificationSetting_EmailSetting) bool {
|
||||
if setting == nil || existing == nil {
|
||||
return false
|
||||
}
|
||||
return strings.TrimSpace(setting.SmtpHost) == strings.TrimSpace(existing.SmtpHost) &&
|
||||
setting.SmtpPort == existing.SmtpPort &&
|
||||
strings.TrimSpace(setting.SmtpUsername) == strings.TrimSpace(existing.SmtpUsername) &&
|
||||
setting.UseTls == existing.UseTls &&
|
||||
setting.UseSsl == existing.UseSsl
|
||||
}
|
||||
|
||||
func convertInstanceSettingFromStore(setting *storepb.InstanceSetting) *v1pb.InstanceSetting {
|
||||
instanceSetting := &v1pb.InstanceSetting{
|
||||
Name: fmt.Sprintf("instance/settings/%s", setting.Key.String()),
|
||||
}
|
||||
switch setting.Value.(type) {
|
||||
case *storepb.InstanceSetting_GeneralSetting:
|
||||
instanceSetting.Value = &v1pb.InstanceSetting_GeneralSetting_{
|
||||
GeneralSetting: convertInstanceGeneralSettingFromStore(setting.GetGeneralSetting()),
|
||||
}
|
||||
case *storepb.InstanceSetting_StorageSetting:
|
||||
instanceSetting.Value = &v1pb.InstanceSetting_StorageSetting_{
|
||||
StorageSetting: convertInstanceStorageSettingFromStore(setting.GetStorageSetting()),
|
||||
}
|
||||
case *storepb.InstanceSetting_MemoRelatedSetting:
|
||||
instanceSetting.Value = &v1pb.InstanceSetting_MemoRelatedSetting_{
|
||||
MemoRelatedSetting: convertInstanceMemoRelatedSettingFromStore(setting.GetMemoRelatedSetting()),
|
||||
}
|
||||
case *storepb.InstanceSetting_TagsSetting:
|
||||
instanceSetting.Value = &v1pb.InstanceSetting_TagsSetting_{
|
||||
TagsSetting: convertInstanceTagsSettingFromStore(setting.GetTagsSetting()),
|
||||
}
|
||||
case *storepb.InstanceSetting_NotificationSetting:
|
||||
instanceSetting.Value = &v1pb.InstanceSetting_NotificationSetting_{
|
||||
NotificationSetting: convertInstanceNotificationSettingFromStore(setting.GetNotificationSetting()),
|
||||
}
|
||||
case *storepb.InstanceSetting_AiSetting:
|
||||
instanceSetting.Value = &v1pb.InstanceSetting_AiSetting{
|
||||
AiSetting: convertInstanceAISettingFromStore(setting.GetAiSetting()),
|
||||
}
|
||||
default:
|
||||
// Leave Value unset for unsupported setting variants.
|
||||
}
|
||||
return instanceSetting
|
||||
}
|
||||
|
||||
func convertInstanceSettingToStore(setting *v1pb.InstanceSetting) *storepb.InstanceSetting {
|
||||
settingKeyString, _ := ExtractInstanceSettingKeyFromName(setting.Name)
|
||||
instanceSetting := &storepb.InstanceSetting{
|
||||
Key: storepb.InstanceSettingKey(storepb.InstanceSettingKey_value[settingKeyString]),
|
||||
Value: &storepb.InstanceSetting_GeneralSetting{
|
||||
GeneralSetting: convertInstanceGeneralSettingToStore(setting.GetGeneralSetting()),
|
||||
},
|
||||
}
|
||||
switch instanceSetting.Key {
|
||||
case storepb.InstanceSettingKey_GENERAL:
|
||||
instanceSetting.Value = &storepb.InstanceSetting_GeneralSetting{
|
||||
GeneralSetting: convertInstanceGeneralSettingToStore(setting.GetGeneralSetting()),
|
||||
}
|
||||
case storepb.InstanceSettingKey_STORAGE:
|
||||
instanceSetting.Value = &storepb.InstanceSetting_StorageSetting{
|
||||
StorageSetting: convertInstanceStorageSettingToStore(setting.GetStorageSetting()),
|
||||
}
|
||||
case storepb.InstanceSettingKey_MEMO_RELATED:
|
||||
instanceSetting.Value = &storepb.InstanceSetting_MemoRelatedSetting{
|
||||
MemoRelatedSetting: convertInstanceMemoRelatedSettingToStore(setting.GetMemoRelatedSetting()),
|
||||
}
|
||||
case storepb.InstanceSettingKey_TAGS:
|
||||
instanceSetting.Value = &storepb.InstanceSetting_TagsSetting{
|
||||
TagsSetting: convertInstanceTagsSettingToStore(setting.GetTagsSetting()),
|
||||
}
|
||||
case storepb.InstanceSettingKey_NOTIFICATION:
|
||||
instanceSetting.Value = &storepb.InstanceSetting_NotificationSetting{
|
||||
NotificationSetting: convertInstanceNotificationSettingToStore(setting.GetNotificationSetting()),
|
||||
}
|
||||
case storepb.InstanceSettingKey_AI:
|
||||
instanceSetting.Value = &storepb.InstanceSetting_AiSetting{
|
||||
AiSetting: convertInstanceAISettingToStore(setting.GetAiSetting()),
|
||||
}
|
||||
default:
|
||||
// Keep the default GeneralSetting value
|
||||
}
|
||||
return instanceSetting
|
||||
}
|
||||
|
||||
func convertInstanceGeneralSettingFromStore(setting *storepb.InstanceGeneralSetting) *v1pb.InstanceSetting_GeneralSetting {
|
||||
if setting == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
generalSetting := &v1pb.InstanceSetting_GeneralSetting{
|
||||
DisallowUserRegistration: setting.DisallowUserRegistration,
|
||||
DisallowPasswordAuth: setting.DisallowPasswordAuth,
|
||||
AdditionalScript: setting.AdditionalScript,
|
||||
AdditionalStyle: setting.AdditionalStyle,
|
||||
WeekStartDayOffset: setting.WeekStartDayOffset,
|
||||
DisallowChangeUsername: setting.DisallowChangeUsername,
|
||||
DisallowChangeNickname: setting.DisallowChangeNickname,
|
||||
}
|
||||
if setting.CustomProfile != nil {
|
||||
generalSetting.CustomProfile = &v1pb.InstanceSetting_GeneralSetting_CustomProfile{
|
||||
Title: setting.CustomProfile.Title,
|
||||
Description: setting.CustomProfile.Description,
|
||||
LogoUrl: setting.CustomProfile.LogoUrl,
|
||||
}
|
||||
}
|
||||
return generalSetting
|
||||
}
|
||||
|
||||
func convertInstanceGeneralSettingToStore(setting *v1pb.InstanceSetting_GeneralSetting) *storepb.InstanceGeneralSetting {
|
||||
if setting == nil {
|
||||
return nil
|
||||
}
|
||||
generalSetting := &storepb.InstanceGeneralSetting{
|
||||
DisallowUserRegistration: setting.DisallowUserRegistration,
|
||||
DisallowPasswordAuth: setting.DisallowPasswordAuth,
|
||||
AdditionalScript: setting.AdditionalScript,
|
||||
AdditionalStyle: setting.AdditionalStyle,
|
||||
WeekStartDayOffset: setting.WeekStartDayOffset,
|
||||
DisallowChangeUsername: setting.DisallowChangeUsername,
|
||||
DisallowChangeNickname: setting.DisallowChangeNickname,
|
||||
}
|
||||
if setting.CustomProfile != nil {
|
||||
generalSetting.CustomProfile = &storepb.InstanceCustomProfile{
|
||||
Title: setting.CustomProfile.Title,
|
||||
Description: setting.CustomProfile.Description,
|
||||
LogoUrl: setting.CustomProfile.LogoUrl,
|
||||
}
|
||||
}
|
||||
return generalSetting
|
||||
}
|
||||
|
||||
func convertInstanceStorageSettingFromStore(settingpb *storepb.InstanceStorageSetting) *v1pb.InstanceSetting_StorageSetting {
|
||||
if settingpb == nil {
|
||||
return nil
|
||||
}
|
||||
setting := &v1pb.InstanceSetting_StorageSetting{
|
||||
StorageType: v1pb.InstanceSetting_StorageSetting_StorageType(settingpb.StorageType),
|
||||
FilepathTemplate: settingpb.FilepathTemplate,
|
||||
UploadSizeLimitMb: settingpb.UploadSizeLimitMb,
|
||||
}
|
||||
if settingpb.S3Config != nil {
|
||||
setting.S3Config = &v1pb.InstanceSetting_StorageSetting_S3Config{
|
||||
AccessKeyId: settingpb.S3Config.AccessKeyId,
|
||||
// AccessKeySecret is write-only: never returned in responses.
|
||||
Endpoint: settingpb.S3Config.Endpoint,
|
||||
Region: settingpb.S3Config.Region,
|
||||
Bucket: settingpb.S3Config.Bucket,
|
||||
UsePathStyle: settingpb.S3Config.UsePathStyle,
|
||||
}
|
||||
}
|
||||
return setting
|
||||
}
|
||||
|
||||
func convertInstanceStorageSettingToStore(setting *v1pb.InstanceSetting_StorageSetting) *storepb.InstanceStorageSetting {
|
||||
if setting == nil {
|
||||
return nil
|
||||
}
|
||||
settingpb := &storepb.InstanceStorageSetting{
|
||||
StorageType: storepb.InstanceStorageSetting_StorageType(setting.StorageType),
|
||||
FilepathTemplate: setting.FilepathTemplate,
|
||||
UploadSizeLimitMb: setting.UploadSizeLimitMb,
|
||||
}
|
||||
if setting.S3Config != nil {
|
||||
settingpb.S3Config = &storepb.StorageS3Config{
|
||||
AccessKeyId: setting.S3Config.AccessKeyId,
|
||||
AccessKeySecret: setting.S3Config.AccessKeySecret,
|
||||
Endpoint: setting.S3Config.Endpoint,
|
||||
Region: setting.S3Config.Region,
|
||||
Bucket: setting.S3Config.Bucket,
|
||||
UsePathStyle: setting.S3Config.UsePathStyle,
|
||||
}
|
||||
}
|
||||
return settingpb
|
||||
}
|
||||
|
||||
func convertInstanceMemoRelatedSettingFromStore(setting *storepb.InstanceMemoRelatedSetting) *v1pb.InstanceSetting_MemoRelatedSetting {
|
||||
if setting == nil {
|
||||
return nil
|
||||
}
|
||||
return &v1pb.InstanceSetting_MemoRelatedSetting{
|
||||
ContentLengthLimit: setting.ContentLengthLimit,
|
||||
EnableDoubleClickEdit: setting.EnableDoubleClickEdit,
|
||||
Reactions: setting.Reactions,
|
||||
}
|
||||
}
|
||||
|
||||
func convertInstanceMemoRelatedSettingToStore(setting *v1pb.InstanceSetting_MemoRelatedSetting) *storepb.InstanceMemoRelatedSetting {
|
||||
if setting == nil {
|
||||
return nil
|
||||
}
|
||||
return &storepb.InstanceMemoRelatedSetting{
|
||||
ContentLengthLimit: setting.ContentLengthLimit,
|
||||
EnableDoubleClickEdit: setting.EnableDoubleClickEdit,
|
||||
Reactions: setting.Reactions,
|
||||
}
|
||||
}
|
||||
|
||||
func convertInstanceTagsSettingFromStore(setting *storepb.InstanceTagsSetting) *v1pb.InstanceSetting_TagsSetting {
|
||||
if setting == nil {
|
||||
return nil
|
||||
}
|
||||
tags := make(map[string]*v1pb.InstanceSetting_TagMetadata, len(setting.Tags))
|
||||
for tag, metadata := range setting.Tags {
|
||||
tags[tag] = &v1pb.InstanceSetting_TagMetadata{
|
||||
BackgroundColor: metadata.GetBackgroundColor(),
|
||||
BlurContent: metadata.GetBlurContent(),
|
||||
}
|
||||
}
|
||||
return &v1pb.InstanceSetting_TagsSetting{
|
||||
Tags: tags,
|
||||
}
|
||||
}
|
||||
|
||||
func convertInstanceTagsSettingToStore(setting *v1pb.InstanceSetting_TagsSetting) *storepb.InstanceTagsSetting {
|
||||
if setting == nil {
|
||||
return nil
|
||||
}
|
||||
tags := make(map[string]*storepb.InstanceTagMetadata, len(setting.Tags))
|
||||
for tag, metadata := range setting.Tags {
|
||||
tags[tag] = &storepb.InstanceTagMetadata{
|
||||
BackgroundColor: metadata.GetBackgroundColor(),
|
||||
BlurContent: metadata.GetBlurContent(),
|
||||
}
|
||||
}
|
||||
return &storepb.InstanceTagsSetting{
|
||||
Tags: tags,
|
||||
}
|
||||
}
|
||||
|
||||
func convertInstanceNotificationSettingFromStore(setting *storepb.InstanceNotificationSetting) *v1pb.InstanceSetting_NotificationSetting {
|
||||
if setting == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
notificationSetting := &v1pb.InstanceSetting_NotificationSetting{}
|
||||
if setting.Email != nil {
|
||||
notificationSetting.Email = &v1pb.InstanceSetting_NotificationSetting_EmailSetting{
|
||||
Enabled: setting.Email.Enabled,
|
||||
SmtpHost: setting.Email.SmtpHost,
|
||||
SmtpPort: setting.Email.SmtpPort,
|
||||
SmtpUsername: setting.Email.SmtpUsername,
|
||||
// SmtpPassword is write-only: never returned in responses.
|
||||
FromEmail: setting.Email.FromEmail,
|
||||
FromName: setting.Email.FromName,
|
||||
ReplyTo: setting.Email.ReplyTo,
|
||||
UseTls: setting.Email.UseTls,
|
||||
UseSsl: setting.Email.UseSsl,
|
||||
}
|
||||
}
|
||||
return notificationSetting
|
||||
}
|
||||
|
||||
func convertInstanceNotificationSettingToStore(setting *v1pb.InstanceSetting_NotificationSetting) *storepb.InstanceNotificationSetting {
|
||||
if setting == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
notificationSetting := &storepb.InstanceNotificationSetting{}
|
||||
if setting.Email != nil {
|
||||
notificationSetting.Email = &storepb.InstanceNotificationSetting_EmailSetting{
|
||||
Enabled: setting.Email.Enabled,
|
||||
SmtpHost: setting.Email.SmtpHost,
|
||||
SmtpPort: setting.Email.SmtpPort,
|
||||
SmtpUsername: setting.Email.SmtpUsername,
|
||||
SmtpPassword: setting.Email.SmtpPassword,
|
||||
FromEmail: setting.Email.FromEmail,
|
||||
FromName: setting.Email.FromName,
|
||||
ReplyTo: setting.Email.ReplyTo,
|
||||
UseTls: setting.Email.UseTls,
|
||||
UseSsl: setting.Email.UseSsl,
|
||||
}
|
||||
}
|
||||
return notificationSetting
|
||||
}
|
||||
|
||||
func convertInstanceAISettingFromStore(setting *storepb.InstanceAISetting) *v1pb.InstanceSetting_AISetting {
|
||||
if setting == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
aiSetting := &v1pb.InstanceSetting_AISetting{
|
||||
Providers: make([]*v1pb.InstanceSetting_AIProviderConfig, 0, len(setting.Providers)),
|
||||
Transcription: convertTranscriptionConfigFromStore(setting.GetTranscription()),
|
||||
}
|
||||
for _, provider := range setting.Providers {
|
||||
if provider == nil {
|
||||
continue
|
||||
}
|
||||
apiKey := provider.GetApiKey()
|
||||
aiSetting.Providers = append(aiSetting.Providers, &v1pb.InstanceSetting_AIProviderConfig{
|
||||
Id: provider.GetId(),
|
||||
Title: provider.GetTitle(),
|
||||
Type: v1pb.InstanceSetting_AIProviderType(provider.GetType()),
|
||||
Endpoint: provider.GetEndpoint(),
|
||||
ApiKeySet: apiKey != "",
|
||||
ApiKeyHint: maskAPIKey(apiKey),
|
||||
})
|
||||
}
|
||||
return aiSetting
|
||||
}
|
||||
|
||||
func convertInstanceAISettingToStore(setting *v1pb.InstanceSetting_AISetting) *storepb.InstanceAISetting {
|
||||
if setting == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
aiSetting := &storepb.InstanceAISetting{
|
||||
Providers: make([]*storepb.AIProviderConfig, 0, len(setting.Providers)),
|
||||
Transcription: convertTranscriptionConfigToStore(setting.GetTranscription()),
|
||||
}
|
||||
for _, provider := range setting.Providers {
|
||||
if provider == nil {
|
||||
continue
|
||||
}
|
||||
aiSetting.Providers = append(aiSetting.Providers, &storepb.AIProviderConfig{
|
||||
Id: provider.GetId(),
|
||||
Title: provider.GetTitle(),
|
||||
Type: storepb.AIProviderType(provider.GetType()),
|
||||
Endpoint: provider.GetEndpoint(),
|
||||
ApiKey: provider.GetApiKey(),
|
||||
})
|
||||
}
|
||||
return aiSetting
|
||||
}
|
||||
|
||||
func convertTranscriptionConfigFromStore(setting *storepb.TranscriptionConfig) *v1pb.InstanceSetting_TranscriptionConfig {
|
||||
if setting == nil {
|
||||
return nil
|
||||
}
|
||||
return &v1pb.InstanceSetting_TranscriptionConfig{
|
||||
ProviderId: setting.GetProviderId(),
|
||||
Model: setting.GetModel(),
|
||||
Language: setting.GetLanguage(),
|
||||
Prompt: setting.GetPrompt(),
|
||||
}
|
||||
}
|
||||
|
||||
func convertTranscriptionConfigToStore(setting *v1pb.InstanceSetting_TranscriptionConfig) *storepb.TranscriptionConfig {
|
||||
if setting == nil {
|
||||
return nil
|
||||
}
|
||||
return &storepb.TranscriptionConfig{
|
||||
ProviderId: setting.GetProviderId(),
|
||||
Model: setting.GetModel(),
|
||||
Language: setting.GetLanguage(),
|
||||
Prompt: setting.GetPrompt(),
|
||||
}
|
||||
}
|
||||
|
||||
func validateInstanceSetting(setting *v1pb.InstanceSetting) error {
|
||||
key, err := ExtractInstanceSettingKeyFromName(setting.Name)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if key != storepb.InstanceSettingKey_TAGS.String() {
|
||||
return nil
|
||||
}
|
||||
return validateInstanceTagsSetting(setting.GetTagsSetting())
|
||||
}
|
||||
|
||||
func (s *APIV1Service) prepareInstanceAISettingForUpdate(ctx context.Context, setting *storepb.InstanceAISetting) error {
|
||||
if setting == nil {
|
||||
return errors.New("AI setting is required")
|
||||
}
|
||||
|
||||
existing, err := s.Store.GetInstanceAISetting(ctx)
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "failed to get existing AI setting")
|
||||
}
|
||||
existingProviders := map[string]*storepb.AIProviderConfig{}
|
||||
if existing != nil {
|
||||
for _, provider := range existing.Providers {
|
||||
if provider != nil && provider.Id != "" {
|
||||
existingProviders[provider.Id] = provider
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
seenIDs := map[string]bool{}
|
||||
for _, provider := range setting.Providers {
|
||||
if provider == nil {
|
||||
return errors.New("provider cannot be nil")
|
||||
}
|
||||
|
||||
provider.Id = strings.TrimSpace(provider.Id)
|
||||
if provider.Id == "" {
|
||||
provider.Id = shortuuid.New()
|
||||
}
|
||||
if seenIDs[provider.Id] {
|
||||
return errors.Errorf("duplicate provider ID %q", provider.Id)
|
||||
}
|
||||
seenIDs[provider.Id] = true
|
||||
|
||||
provider.Title = strings.TrimSpace(provider.Title)
|
||||
if provider.Title == "" {
|
||||
return errors.New("provider title is required")
|
||||
}
|
||||
if provider.Type != storepb.AIProviderType_OPENAI && provider.Type != storepb.AIProviderType_GEMINI {
|
||||
return errors.Errorf("provider %q has unsupported type", provider.Id)
|
||||
}
|
||||
|
||||
provider.Endpoint = strings.TrimSpace(provider.Endpoint)
|
||||
if provider.Type == storepb.AIProviderType_OPENAI && provider.Endpoint == "" {
|
||||
provider.Endpoint = "https://api.openai.com/v1"
|
||||
}
|
||||
if provider.Type == storepb.AIProviderType_GEMINI && provider.Endpoint == "" {
|
||||
provider.Endpoint = "https://generativelanguage.googleapis.com/v1beta"
|
||||
}
|
||||
|
||||
if provider.ApiKey == "" {
|
||||
if existingProvider, ok := existingProviders[provider.Id]; ok {
|
||||
provider.ApiKey = existingProvider.ApiKey
|
||||
}
|
||||
}
|
||||
if provider.ApiKey == "" {
|
||||
return errors.Errorf("provider %q API key is required", provider.Id)
|
||||
}
|
||||
}
|
||||
|
||||
if err := preparePersistedTranscriptionConfig(setting, existing); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func preparePersistedTranscriptionConfig(setting *storepb.InstanceAISetting, existing *storepb.InstanceAISetting) error {
|
||||
// Preserve the previously stored transcription config when the request omits it,
|
||||
// matching the same "absence == keep" semantics used for API keys. The preserved
|
||||
// config still falls through to validation below, so a stale provider_id is
|
||||
// rejected if the same update removed or renamed its referenced provider.
|
||||
if setting.Transcription == nil && existing != nil {
|
||||
setting.Transcription = existing.GetTranscription()
|
||||
}
|
||||
if setting.Transcription == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
cfg := setting.Transcription
|
||||
cfg.ProviderId = strings.TrimSpace(cfg.ProviderId)
|
||||
cfg.Model = strings.TrimSpace(cfg.Model)
|
||||
cfg.Language = strings.TrimSpace(cfg.Language)
|
||||
cfg.Prompt = strings.TrimSpace(cfg.Prompt)
|
||||
|
||||
if cfg.ProviderId != "" {
|
||||
referenced := false
|
||||
for _, provider := range setting.Providers {
|
||||
if provider != nil && provider.Id == cfg.ProviderId {
|
||||
referenced = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !referenced {
|
||||
return errors.Errorf("transcription provider_id %q does not reference any configured provider", cfg.ProviderId)
|
||||
}
|
||||
}
|
||||
|
||||
if len(cfg.Model) > maxTranscriptionConfigModelLength {
|
||||
return errors.Errorf("transcription model is too long; maximum length is %d characters", maxTranscriptionConfigModelLength)
|
||||
}
|
||||
if len(cfg.Language) > maxTranscriptionConfigLanguageLength {
|
||||
return errors.Errorf("transcription language is too long; maximum length is %d characters", maxTranscriptionConfigLanguageLength)
|
||||
}
|
||||
if len(cfg.Prompt) > maxTranscriptionConfigPromptLength {
|
||||
return errors.Errorf("transcription prompt is too long; maximum length is %d characters", maxTranscriptionConfigPromptLength)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func maskAPIKey(apiKey string) string {
|
||||
if apiKey == "" {
|
||||
return ""
|
||||
}
|
||||
if len(apiKey) <= 8 {
|
||||
return "..."
|
||||
}
|
||||
prefixLength := min(4, len(apiKey))
|
||||
return apiKey[:prefixLength] + "..." + apiKey[len(apiKey)-4:]
|
||||
}
|
||||
|
||||
func validateInstanceTagsSetting(setting *v1pb.InstanceSetting_TagsSetting) error {
|
||||
if setting == nil {
|
||||
return errors.New("tags setting is required")
|
||||
}
|
||||
for tag, metadata := range setting.Tags {
|
||||
if strings.TrimSpace(tag) == "" {
|
||||
return errors.New("tag key cannot be empty")
|
||||
}
|
||||
if _, err := regexp.Compile(tag); err != nil {
|
||||
return errors.Errorf("tag key %q is not a valid regex pattern: %v", tag, err)
|
||||
}
|
||||
if metadata == nil {
|
||||
return errors.Errorf("tag metadata is required for %q", tag)
|
||||
}
|
||||
if metadata.GetBackgroundColor() != nil {
|
||||
if err := validateInstanceColor(metadata.GetBackgroundColor()); err != nil {
|
||||
return errors.Wrapf(err, "background_color for %q", tag)
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateInstanceColor(color *colorpb.Color) error {
|
||||
if err := validateInstanceColorComponent("red", color.GetRed()); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateInstanceColorComponent("green", color.GetGreen()); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateInstanceColorComponent("blue", color.GetBlue()); err != nil {
|
||||
return err
|
||||
}
|
||||
if alpha := color.GetAlpha(); alpha != nil {
|
||||
if err := validateInstanceColorComponent("alpha", alpha.GetValue()); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateInstanceColorComponent(name string, value float32) error {
|
||||
if math.IsNaN(float64(value)) || math.IsInf(float64(value), 0) {
|
||||
return errors.Errorf("%s must be a finite number", name)
|
||||
}
|
||||
if value < 0 || value > 1 {
|
||||
return errors.Errorf("%s must be between 0 and 1", name)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) GetInstanceAdmin(ctx context.Context) (*v1pb.User, error) {
|
||||
adminUserType := store.RoleAdmin
|
||||
user, err := s.Store.GetUser(ctx, &store.FindUser{
|
||||
Role: &adminUserType,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, errors.Wrapf(err, "failed to find admin")
|
||||
}
|
||||
if user == nil {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
currentUser, _ := s.fetchCurrentUser(ctx)
|
||||
return convertUserFromStore(user, currentUser), nil
|
||||
}
|
||||
@@ -0,0 +1,165 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
"golang.org/x/sync/errgroup"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
"google.golang.org/protobuf/types/known/timestamppb"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
const instanceStatsCacheTTL = 60 * time.Second
|
||||
|
||||
// instanceStatsCache is a single-value, mutex-guarded cache for InstanceStats.
|
||||
type instanceStatsCache struct {
|
||||
mu sync.Mutex
|
||||
value *v1pb.InstanceStats
|
||||
expiry time.Time
|
||||
}
|
||||
|
||||
func (c *instanceStatsCache) get() (*v1pb.InstanceStats, bool) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.value == nil || time.Now().After(c.expiry) {
|
||||
return nil, false
|
||||
}
|
||||
return c.value, true
|
||||
}
|
||||
|
||||
func (c *instanceStatsCache) set(v *v1pb.InstanceStats, ttl time.Duration) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.value = v
|
||||
c.expiry = time.Now().Add(ttl)
|
||||
}
|
||||
|
||||
// GetInstanceStats returns resource usage statistics. Admin only.
|
||||
func (s *APIV1Service) GetInstanceStats(ctx context.Context, _ *v1pb.GetInstanceStatsRequest) (*v1pb.InstanceStats, error) {
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
if user.Role != store.RoleAdmin {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
|
||||
if cached, ok := s.instanceStatsCache.get(); ok {
|
||||
return cached, nil
|
||||
}
|
||||
|
||||
stats, err := s.computeInstanceStats(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to compute instance stats: %v", err)
|
||||
}
|
||||
s.instanceStatsCache.set(stats, instanceStatsCacheTTL)
|
||||
return stats, nil
|
||||
}
|
||||
|
||||
// computeInstanceStats runs all stat subqueries in parallel and assembles the result.
|
||||
// Per-subtask failures degrade to -1 sentinel values; only a total failure (every
|
||||
// subtask errored) is propagated as an error.
|
||||
func (s *APIV1Service) computeInstanceStats(ctx context.Context) (*v1pb.InstanceStats, error) {
|
||||
stats := &v1pb.InstanceStats{
|
||||
Database: &v1pb.InstanceStats_DatabaseStats{
|
||||
Driver: s.Profile.Driver,
|
||||
SizeBytes: -1,
|
||||
},
|
||||
LocalStorageBytes: -1,
|
||||
GeneratedTime: timestamppb.Now(),
|
||||
}
|
||||
|
||||
type result struct {
|
||||
name string
|
||||
err error
|
||||
}
|
||||
var (
|
||||
mu sync.Mutex
|
||||
results []result
|
||||
record = func(name string, err error) {
|
||||
mu.Lock()
|
||||
results = append(results, result{name, err})
|
||||
mu.Unlock()
|
||||
}
|
||||
)
|
||||
|
||||
g, gctx := errgroup.WithContext(ctx)
|
||||
|
||||
g.Go(func() error {
|
||||
size, err := s.Store.GetDriver().GetDatabaseSize(gctx)
|
||||
if err != nil {
|
||||
record("database_size", err)
|
||||
return nil
|
||||
}
|
||||
stats.Database.SizeBytes = size
|
||||
return nil
|
||||
})
|
||||
|
||||
g.Go(func() error {
|
||||
size, err := walkLocalStorage(s.Profile.Data)
|
||||
if err != nil {
|
||||
record("local_storage", err)
|
||||
return nil
|
||||
}
|
||||
stats.LocalStorageBytes = size
|
||||
return nil
|
||||
})
|
||||
|
||||
_ = g.Wait()
|
||||
|
||||
for _, r := range results {
|
||||
slog.Warn("instance stats subtask failed", slog.String("subtask", r.name), slog.String("err", r.err.Error()))
|
||||
}
|
||||
|
||||
const totalSubtasks = 2
|
||||
if len(results) == totalSubtasks {
|
||||
return nil, errors.New("all instance stats subtasks failed")
|
||||
}
|
||||
return stats, nil
|
||||
}
|
||||
|
||||
// walkLocalStorage returns the recursive size of dir in bytes.
|
||||
// Symlinks are not followed; per-entry errors below the root are ignored
|
||||
// (the walk continues). An error accessing the root itself is returned.
|
||||
func walkLocalStorage(dir string) (int64, error) {
|
||||
if dir == "" {
|
||||
return -1, errors.New("empty data directory")
|
||||
}
|
||||
var total int64
|
||||
err := filepath.WalkDir(dir, func(path string, entry os.DirEntry, walkErr error) error {
|
||||
if walkErr != nil {
|
||||
if path == dir {
|
||||
// Root itself is inaccessible — abort the walk.
|
||||
return walkErr
|
||||
}
|
||||
// Ignore per-entry errors (e.g. permission denied on a single file).
|
||||
return nil
|
||||
}
|
||||
if entry.IsDir() {
|
||||
return nil
|
||||
}
|
||||
info, err := entry.Info()
|
||||
if err != nil {
|
||||
// Ignore stat errors on individual entries; continue the walk.
|
||||
return nil //nolint:nilerr
|
||||
}
|
||||
total += info.Size()
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return -1, errors.Wrap(err, "walk failed")
|
||||
}
|
||||
return total, nil
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestWalkLocalStorage_SumsFileSizes(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
require.NoError(t, os.WriteFile(filepath.Join(dir, "a.txt"), []byte("hello"), 0o600)) // 5
|
||||
require.NoError(t, os.WriteFile(filepath.Join(dir, "b.txt"), []byte("world!"), 0o600)) // 6
|
||||
sub := filepath.Join(dir, "sub")
|
||||
require.NoError(t, os.Mkdir(sub, 0o700))
|
||||
require.NoError(t, os.WriteFile(filepath.Join(sub, "c.txt"), []byte("xx"), 0o600)) // 2
|
||||
|
||||
size, err := walkLocalStorage(dir)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(13), size)
|
||||
}
|
||||
|
||||
func TestWalkLocalStorage_EmptyDir(t *testing.T) {
|
||||
size, err := walkLocalStorage("")
|
||||
require.Error(t, err)
|
||||
require.Equal(t, int64(-1), size)
|
||||
}
|
||||
|
||||
func TestWalkLocalStorage_NonexistentDir(t *testing.T) {
|
||||
size, err := walkLocalStorage(filepath.Join(t.TempDir(), "does-not-exist"))
|
||||
require.Error(t, err)
|
||||
require.Equal(t, int64(-1), size)
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
"github.com/usememos/memos/internal/httpgetter"
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
)
|
||||
|
||||
func TestGetLinkMetadata(t *testing.T) {
|
||||
originalFetchHTMLMeta := fetchHTMLMeta
|
||||
t.Cleanup(func() {
|
||||
fetchHTMLMeta = originalFetchHTMLMeta
|
||||
})
|
||||
|
||||
fetchHTMLMeta = func(url string) (*httpgetter.HTMLMeta, error) {
|
||||
require.Equal(t, "https://example.com/article", url)
|
||||
return &httpgetter.HTMLMeta{
|
||||
Title: "Example title",
|
||||
Description: "Example description",
|
||||
Image: "https://example.com/cover.png",
|
||||
}, nil
|
||||
}
|
||||
|
||||
metadata, err := (&APIV1Service{}).GetLinkMetadata(context.Background(), &v1pb.GetLinkMetadataRequest{
|
||||
Url: "https://example.com/article",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "https://example.com/article", metadata.Url)
|
||||
require.Equal(t, "Example title", metadata.Title)
|
||||
require.Equal(t, "Example description", metadata.Description)
|
||||
require.Equal(t, "https://example.com/cover.png", metadata.Image)
|
||||
}
|
||||
|
||||
func TestGetLinkMetadataEmptyURL(t *testing.T) {
|
||||
_, err := (&APIV1Service{}).GetLinkMetadata(context.Background(), &v1pb.GetLinkMetadataRequest{})
|
||||
require.Error(t, err)
|
||||
require.Equal(t, codes.InvalidArgument, status.Code(err))
|
||||
}
|
||||
|
||||
func TestGetLinkMetadataInternalURL(t *testing.T) {
|
||||
_, err := (&APIV1Service{}).GetLinkMetadata(context.Background(), &v1pb.GetLinkMetadataRequest{
|
||||
Url: "http://192.168.0.1",
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Equal(t, codes.InvalidArgument, status.Code(err))
|
||||
}
|
||||
|
||||
func TestBatchGetLinkMetadata(t *testing.T) {
|
||||
originalFetchHTMLMeta := fetchHTMLMeta
|
||||
t.Cleanup(func() {
|
||||
fetchHTMLMeta = originalFetchHTMLMeta
|
||||
})
|
||||
|
||||
var fetchedURLs []string
|
||||
fetchHTMLMeta = func(url string) (*httpgetter.HTMLMeta, error) {
|
||||
fetchedURLs = append(fetchedURLs, url)
|
||||
return &httpgetter.HTMLMeta{
|
||||
Title: fmt.Sprintf("Title for %s", url),
|
||||
Description: fmt.Sprintf("Description for %s", url),
|
||||
Image: fmt.Sprintf("%s/cover.png", url),
|
||||
}, nil
|
||||
}
|
||||
|
||||
response, err := (&APIV1Service{}).BatchGetLinkMetadata(context.Background(), &v1pb.BatchGetLinkMetadataRequest{
|
||||
Urls: []string{
|
||||
"https://example.com/one",
|
||||
"https://example.com/two",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []string{"https://example.com/one", "https://example.com/two"}, fetchedURLs)
|
||||
require.Len(t, response.LinkMetadata, 2)
|
||||
require.Equal(t, "https://example.com/one", response.LinkMetadata[0].Url)
|
||||
require.Equal(t, "Title for https://example.com/one", response.LinkMetadata[0].Title)
|
||||
require.Equal(t, "https://example.com/two", response.LinkMetadata[1].Url)
|
||||
require.Equal(t, "Title for https://example.com/two", response.LinkMetadata[1].Title)
|
||||
}
|
||||
|
||||
func TestBatchGetLinkMetadataEmptyURLs(t *testing.T) {
|
||||
_, err := (&APIV1Service{}).BatchGetLinkMetadata(context.Background(), &v1pb.BatchGetLinkMetadataRequest{})
|
||||
require.Error(t, err)
|
||||
require.Equal(t, codes.InvalidArgument, status.Code(err))
|
||||
}
|
||||
|
||||
func TestBatchGetLinkMetadataTooManyURLs(t *testing.T) {
|
||||
urls := make([]string, maxBatchGetLinkMetadata+1)
|
||||
for i := range urls {
|
||||
urls[i] = fmt.Sprintf("https://example.com/%d", i)
|
||||
}
|
||||
|
||||
_, err := (&APIV1Service{}).BatchGetLinkMetadata(context.Background(), &v1pb.BatchGetLinkMetadataRequest{
|
||||
Urls: urls,
|
||||
})
|
||||
require.Error(t, err)
|
||||
require.Equal(t, codes.InvalidArgument, status.Code(err))
|
||||
}
|
||||
@@ -0,0 +1,241 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"slices"
|
||||
"time"
|
||||
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
"google.golang.org/protobuf/types/known/emptypb"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func (s *APIV1Service) SetMemoAttachments(ctx context.Context, request *v1pb.SetMemoAttachmentsRequest) (*emptypb.Empty, error) {
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
memoUID, err := ExtractMemoUIDFromName(request.Name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid memo name: %v", err)
|
||||
}
|
||||
memo, err := s.Store.GetMemo(ctx, &store.FindMemo{UID: &memoUID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get memo")
|
||||
}
|
||||
if memo == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "memo not found")
|
||||
}
|
||||
if !canModifyMemo(user, memo) {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
if err := s.setMemoAttachmentsInternal(ctx, user, memo, request.Attachments); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := s.touchMemoUpdatedTimestamp(ctx, memo.ID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
updatedMemo, parentMemo, memoMessage, err := s.buildUpdatedMemoState(ctx, memo.ID)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to build updated memo state")
|
||||
}
|
||||
s.dispatchMemoUpdatedSideEffects(ctx, updatedMemo, parentMemo, memoMessage)
|
||||
|
||||
return &emptypb.Empty{}, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) setMemoAttachmentsInternal(ctx context.Context, user *store.User, memo *store.Memo, requestAttachments []*v1pb.Attachment) error {
|
||||
currentAttachments, err := s.Store.ListAttachments(ctx, &store.FindAttachment{
|
||||
MemoID: &memo.ID,
|
||||
})
|
||||
if err != nil {
|
||||
return status.Errorf(codes.Internal, "failed to list attachments")
|
||||
}
|
||||
|
||||
normalizedAttachments, err := s.normalizeMemoAttachmentRequest(ctx, user, currentAttachments, requestAttachments)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
requestedIDs := make(map[int32]bool, len(normalizedAttachments))
|
||||
for _, attachment := range normalizedAttachments {
|
||||
requestedIDs[attachment.ID] = true
|
||||
}
|
||||
|
||||
// Delete attachments that are not in the request.
|
||||
for _, attachment := range currentAttachments {
|
||||
if !requestedIDs[attachment.ID] {
|
||||
if attachment.CreatorID != user.ID && !isSuperUser(user) {
|
||||
return status.Errorf(codes.PermissionDenied, "cannot remove another user's attachment")
|
||||
}
|
||||
if err = s.Store.DeleteAttachment(ctx, &store.DeleteAttachment{
|
||||
ID: int32(attachment.ID),
|
||||
MemoID: &memo.ID,
|
||||
}); err != nil {
|
||||
return status.Errorf(codes.Internal, "failed to delete attachment")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
slices.Reverse(normalizedAttachments)
|
||||
// Update attachments' memo_id in the request.
|
||||
for index, attachment := range normalizedAttachments {
|
||||
updatedTs := time.Now().Unix() + int64(index)
|
||||
if err := s.Store.UpdateAttachment(ctx, &store.UpdateAttachment{
|
||||
ID: attachment.ID,
|
||||
MemoID: &memo.ID,
|
||||
UpdatedTs: &updatedTs,
|
||||
}); err != nil {
|
||||
return status.Errorf(codes.Internal, "failed to update attachment: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) normalizeMemoAttachmentRequest(
|
||||
ctx context.Context,
|
||||
user *store.User,
|
||||
currentAttachments []*store.Attachment,
|
||||
requestAttachments []*v1pb.Attachment,
|
||||
) ([]*store.Attachment, error) {
|
||||
requestedAttachments := make([]*store.Attachment, 0, len(requestAttachments))
|
||||
for _, requestAttachment := range requestAttachments {
|
||||
attachmentUID, err := ExtractAttachmentUIDFromName(requestAttachment.Name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid attachment name: %v", err)
|
||||
}
|
||||
attachment, err := s.Store.GetAttachment(ctx, &store.FindAttachment{UID: &attachmentUID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get attachment: %v", err)
|
||||
}
|
||||
if attachment == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "attachment not found: %s", attachmentUID)
|
||||
}
|
||||
if attachment.CreatorID != user.ID && !isSuperUser(user) {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "cannot attach another user's attachment")
|
||||
}
|
||||
requestedAttachments = append(requestedAttachments, attachment)
|
||||
}
|
||||
|
||||
currentGroups := make(map[string][]*store.Attachment)
|
||||
for _, attachment := range currentAttachments {
|
||||
motion := getAttachmentMotionMedia(attachment)
|
||||
if motion == nil || motion.GroupId == "" {
|
||||
continue
|
||||
}
|
||||
currentGroups[motion.GroupId] = append(currentGroups[motion.GroupId], attachment)
|
||||
}
|
||||
|
||||
requestGroups := make(map[string][]*store.Attachment)
|
||||
requestNamesByGroup := make(map[string]map[string]bool)
|
||||
for _, attachment := range requestedAttachments {
|
||||
motion := getAttachmentMotionMedia(attachment)
|
||||
if motion == nil || motion.GroupId == "" {
|
||||
continue
|
||||
}
|
||||
requestGroups[motion.GroupId] = append(requestGroups[motion.GroupId], attachment)
|
||||
if requestNamesByGroup[motion.GroupId] == nil {
|
||||
requestNamesByGroup[motion.GroupId] = make(map[string]bool)
|
||||
}
|
||||
requestNamesByGroup[motion.GroupId][attachment.UID] = true
|
||||
}
|
||||
|
||||
normalized := make([]*store.Attachment, 0, len(requestedAttachments))
|
||||
appendedGroups := make(map[string]bool)
|
||||
appendedAttachments := make(map[string]bool)
|
||||
for _, attachment := range requestedAttachments {
|
||||
motion := getAttachmentMotionMedia(attachment)
|
||||
if motion == nil || motion.GroupId == "" {
|
||||
if !appendedAttachments[attachment.UID] {
|
||||
normalized = append(normalized, attachment)
|
||||
appendedAttachments[attachment.UID] = true
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
groupID := motion.GroupId
|
||||
if appendedGroups[groupID] {
|
||||
continue
|
||||
}
|
||||
|
||||
currentGroup := currentGroups[groupID]
|
||||
if isMultiMemberMotionGroup(currentGroup) && !allGroupMembersRequested(currentGroup, requestNamesByGroup[groupID]) {
|
||||
appendedGroups[groupID] = true
|
||||
continue
|
||||
}
|
||||
|
||||
for _, groupAttachment := range requestGroups[groupID] {
|
||||
if appendedAttachments[groupAttachment.UID] {
|
||||
continue
|
||||
}
|
||||
normalized = append(normalized, groupAttachment)
|
||||
appendedAttachments[groupAttachment.UID] = true
|
||||
}
|
||||
appendedGroups[groupID] = true
|
||||
}
|
||||
|
||||
return normalized, nil
|
||||
}
|
||||
|
||||
func allGroupMembersRequested(group []*store.Attachment, requestedNames map[string]bool) bool {
|
||||
if len(group) == 0 {
|
||||
return false
|
||||
}
|
||||
|
||||
for _, attachment := range group {
|
||||
if !requestedNames[attachment.UID] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (s *APIV1Service) ListMemoAttachments(ctx context.Context, request *v1pb.ListMemoAttachmentsRequest) (*v1pb.ListMemoAttachmentsResponse, error) {
|
||||
memoUID, err := ExtractMemoUIDFromName(request.Name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid memo name: %v", err)
|
||||
}
|
||||
memo, err := s.Store.GetMemo(ctx, &store.FindMemo{UID: &memoUID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get memo: %v", err)
|
||||
}
|
||||
if memo == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "memo not found")
|
||||
}
|
||||
|
||||
// Check memo visibility.
|
||||
if memo.Visibility != store.Public {
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
if memo.Visibility == store.Private && memo.CreatorID != user.ID && !isSuperUser(user) {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
}
|
||||
|
||||
attachments, err := s.Store.ListAttachments(ctx, &store.FindAttachment{
|
||||
MemoID: &memo.ID,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to list attachments: %v", err)
|
||||
}
|
||||
|
||||
response := &v1pb.ListMemoAttachmentsResponse{
|
||||
Attachments: []*v1pb.Attachment{},
|
||||
}
|
||||
for _, attachment := range attachments {
|
||||
response.Attachments = append(response.Attachments, convertAttachmentFromStore(attachment))
|
||||
}
|
||||
return response, nil
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func (s *APIV1Service) listMemosByID(ctx context.Context, memoIDs []int32) (map[int32]*store.Memo, error) {
|
||||
if len(memoIDs) == 0 {
|
||||
return map[int32]*store.Memo{}, nil
|
||||
}
|
||||
|
||||
uniqueMemoIDs := make([]int32, 0, len(memoIDs))
|
||||
seenMemoIDs := make(map[int32]struct{}, len(memoIDs))
|
||||
for _, memoID := range memoIDs {
|
||||
if _, seen := seenMemoIDs[memoID]; seen {
|
||||
continue
|
||||
}
|
||||
seenMemoIDs[memoID] = struct{}{}
|
||||
uniqueMemoIDs = append(uniqueMemoIDs, memoID)
|
||||
}
|
||||
|
||||
memos, err := s.Store.ListMemos(ctx, &store.FindMemo{IDList: uniqueMemoIDs})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
memosByID := make(map[int32]*store.Memo, len(memos))
|
||||
for _, memo := range memos {
|
||||
memosByID[memo.ID] = memo
|
||||
}
|
||||
return memosByID, nil
|
||||
}
|
||||
@@ -0,0 +1,146 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
// suppressMentionKey is a context key used to suppress mention notification side effects
|
||||
// when CreateMemo is called internally from CreateMemoComment.
|
||||
type suppressMentionKey struct{}
|
||||
|
||||
func withSuppressMentionNotifications(ctx context.Context) context.Context {
|
||||
return context.WithValue(ctx, suppressMentionKey{}, true)
|
||||
}
|
||||
|
||||
func isMentionNotificationSuppressed(ctx context.Context) bool {
|
||||
v, ok := ctx.Value(suppressMentionKey{}).(bool)
|
||||
return ok && v
|
||||
}
|
||||
|
||||
func (s *APIV1Service) resolveMentionTargets(ctx context.Context, content string) (map[int32]*store.User, error) {
|
||||
targets := make(map[int32]*store.User)
|
||||
if content == "" {
|
||||
return targets, nil
|
||||
}
|
||||
|
||||
data, err := s.MarkdownService.ExtractAll([]byte(content))
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to extract mentions")
|
||||
}
|
||||
if len(data.Mentions) == 0 {
|
||||
return targets, nil
|
||||
}
|
||||
|
||||
normal := store.Normal
|
||||
users, err := s.Store.ListUsers(ctx, &store.FindUser{
|
||||
UsernameList: data.Mentions,
|
||||
RowStatus: &normal,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to resolve mention users")
|
||||
}
|
||||
|
||||
for _, user := range users {
|
||||
targets[user.ID] = user
|
||||
}
|
||||
|
||||
return targets, nil
|
||||
}
|
||||
|
||||
func canUserAccessMentionContext(target *store.User, memo *store.Memo, relatedMemo *store.Memo) bool {
|
||||
if target == nil || memo == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
if relatedMemo != nil {
|
||||
if relatedMemo.Visibility == store.Private && target.ID != relatedMemo.CreatorID {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
if memo.Visibility == store.Private && target.ID != memo.CreatorID {
|
||||
return false
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
func shouldSkipMentionInbox(target *store.User, memo *store.Memo, relatedMemo *store.Memo) bool {
|
||||
if target == nil || memo == nil {
|
||||
return true
|
||||
}
|
||||
|
||||
if target.ID == memo.CreatorID {
|
||||
return true
|
||||
}
|
||||
|
||||
// Comment creation already generates a memo-comment inbox item for the parent creator.
|
||||
if relatedMemo != nil && target.ID == relatedMemo.CreatorID && memo.Visibility != store.Private && memo.CreatorID != relatedMemo.CreatorID {
|
||||
return true
|
||||
}
|
||||
|
||||
return !canUserAccessMentionContext(target, memo, relatedMemo)
|
||||
}
|
||||
|
||||
func (s *APIV1Service) dispatchMemoMentionNotifications(ctx context.Context, memo *store.Memo, relatedMemo *store.Memo, previousContent string) error {
|
||||
if memo == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
currentTargets, err := s.resolveMentionTargets(ctx, memo.Content)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if len(currentTargets) == 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
previousTargets, err := s.resolveMentionTargets(ctx, previousContent)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for userID, target := range currentTargets {
|
||||
if _, exists := previousTargets[userID]; exists {
|
||||
continue
|
||||
}
|
||||
if shouldSkipMentionInbox(target, memo, relatedMemo) {
|
||||
continue
|
||||
}
|
||||
|
||||
payload := &storepb.InboxMessage_MemoMentionPayload{
|
||||
MemoId: memo.ID,
|
||||
}
|
||||
if relatedMemo != nil {
|
||||
payload.RelatedMemoId = relatedMemo.ID
|
||||
}
|
||||
|
||||
if _, err := s.createInboxWithEmailNotification(ctx, &store.Inbox{
|
||||
SenderID: memo.CreatorID,
|
||||
ReceiverID: target.ID,
|
||||
Status: store.UNREAD,
|
||||
Message: &storepb.InboxMessage{
|
||||
Type: storepb.InboxMessage_MEMO_MENTION,
|
||||
Payload: &storepb.InboxMessage_MemoMention{
|
||||
MemoMention: payload,
|
||||
},
|
||||
},
|
||||
}); err != nil {
|
||||
return errors.Wrap(err, "failed to create mention inbox")
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) dispatchMemoMentionNotificationsBestEffort(ctx context.Context, memo *store.Memo, relatedMemo *store.Memo, previousContent string) {
|
||||
if err := s.dispatchMemoMentionNotifications(ctx, memo, relatedMemo, previousContent); err != nil {
|
||||
slog.Warn("Failed to dispatch memo mention notifications", slog.Any("err", err), slog.Int64("memo_id", int64(memo.ID)))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,197 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
"google.golang.org/protobuf/types/known/emptypb"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func (s *APIV1Service) SetMemoRelations(ctx context.Context, request *v1pb.SetMemoRelationsRequest) (*emptypb.Empty, error) {
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
memoUID, err := ExtractMemoUIDFromName(request.Name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid memo name: %v", err)
|
||||
}
|
||||
memo, err := s.Store.GetMemo(ctx, &store.FindMemo{UID: &memoUID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get memo")
|
||||
}
|
||||
if memo == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "memo not found")
|
||||
}
|
||||
if memo.CreatorID != user.ID && !isSuperUser(user) {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
if err := s.setMemoRelationsInternal(ctx, memo, request.Relations); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := s.touchMemoUpdatedTimestamp(ctx, memo.ID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
updatedMemo, parentMemo, memoMessage, err := s.buildUpdatedMemoState(ctx, memo.ID)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to build updated memo state")
|
||||
}
|
||||
s.dispatchMemoUpdatedSideEffects(ctx, updatedMemo, parentMemo, memoMessage)
|
||||
|
||||
return &emptypb.Empty{}, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) setMemoRelationsInternal(ctx context.Context, memo *store.Memo, relations []*v1pb.MemoRelation) error {
|
||||
referenceType := store.MemoRelationReference
|
||||
// Delete all reference relations first.
|
||||
if err := s.Store.DeleteMemoRelation(ctx, &store.DeleteMemoRelation{
|
||||
MemoID: &memo.ID,
|
||||
Type: &referenceType,
|
||||
}); err != nil {
|
||||
return status.Errorf(codes.Internal, "failed to delete memo relation")
|
||||
}
|
||||
|
||||
for _, relation := range relations {
|
||||
// Ignore reflexive relations.
|
||||
if buildMemoName(memo.UID) == relation.RelatedMemo.Name {
|
||||
continue
|
||||
}
|
||||
// Ignore comment relations as there's no need to update a comment's relation.
|
||||
// Inserting/Deleting a comment is handled elsewhere.
|
||||
if relation.Type == v1pb.MemoRelation_COMMENT {
|
||||
continue
|
||||
}
|
||||
relatedMemoUID, err := ExtractMemoUIDFromName(relation.RelatedMemo.Name)
|
||||
if err != nil {
|
||||
return status.Errorf(codes.InvalidArgument, "invalid related memo name: %v", err)
|
||||
}
|
||||
relatedMemo, err := s.Store.GetMemo(ctx, &store.FindMemo{UID: &relatedMemoUID})
|
||||
if err != nil {
|
||||
return status.Errorf(codes.Internal, "failed to get related memo")
|
||||
}
|
||||
if _, err := s.Store.UpsertMemoRelation(ctx, &store.MemoRelation{
|
||||
MemoID: memo.ID,
|
||||
RelatedMemoID: relatedMemo.ID,
|
||||
Type: convertMemoRelationTypeToStore(relation.Type),
|
||||
}); err != nil {
|
||||
return status.Errorf(codes.Internal, "failed to upsert memo relation")
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) ListMemoRelations(ctx context.Context, request *v1pb.ListMemoRelationsRequest) (*v1pb.ListMemoRelationsResponse, error) {
|
||||
memoUID, err := ExtractMemoUIDFromName(request.Name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid memo name: %v", err)
|
||||
}
|
||||
memo, err := s.Store.GetMemo(ctx, &store.FindMemo{UID: &memoUID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get memo")
|
||||
}
|
||||
|
||||
currentUser, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get user")
|
||||
}
|
||||
var memoFilter string
|
||||
if currentUser == nil {
|
||||
memoFilter = `visibility == "PUBLIC"`
|
||||
} else {
|
||||
memoFilter = fmt.Sprintf(`creator_id == %d || visibility in ["PUBLIC", "PROTECTED"]`, currentUser.ID)
|
||||
}
|
||||
relationList := []*v1pb.MemoRelation{}
|
||||
tempList, err := s.Store.ListMemoRelations(ctx, &store.FindMemoRelation{
|
||||
MemoID: &memo.ID,
|
||||
MemoFilter: &memoFilter,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to list memo relations: %v", err)
|
||||
}
|
||||
for _, raw := range tempList {
|
||||
relation, err := s.convertMemoRelationFromStore(ctx, raw)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to convert memo relation")
|
||||
}
|
||||
relationList = append(relationList, relation)
|
||||
}
|
||||
tempList, err = s.Store.ListMemoRelations(ctx, &store.FindMemoRelation{
|
||||
RelatedMemoID: &memo.ID,
|
||||
MemoFilter: &memoFilter,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to list related memo relations: %v", err)
|
||||
}
|
||||
for _, raw := range tempList {
|
||||
relation, err := s.convertMemoRelationFromStore(ctx, raw)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to convert memo relation")
|
||||
}
|
||||
relationList = append(relationList, relation)
|
||||
}
|
||||
|
||||
response := &v1pb.ListMemoRelationsResponse{
|
||||
Relations: relationList,
|
||||
}
|
||||
return response, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) convertMemoRelationFromStore(ctx context.Context, memoRelation *store.MemoRelation) (*v1pb.MemoRelation, error) {
|
||||
memo, err := s.Store.GetMemo(ctx, &store.FindMemo{ID: &memoRelation.MemoID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get memo: %v", err)
|
||||
}
|
||||
memoSnippet, err := s.getMemoContentSnippet(memo.Content)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to get memo content snippet")
|
||||
}
|
||||
relatedMemo, err := s.Store.GetMemo(ctx, &store.FindMemo{ID: &memoRelation.RelatedMemoID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get related memo: %v", err)
|
||||
}
|
||||
relatedMemoSnippet, err := s.getMemoContentSnippet(relatedMemo.Content)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to get related memo content snippet")
|
||||
}
|
||||
return &v1pb.MemoRelation{
|
||||
Memo: &v1pb.MemoRelation_Memo{
|
||||
Name: fmt.Sprintf("%s%s", MemoNamePrefix, memo.UID),
|
||||
Snippet: memoSnippet,
|
||||
},
|
||||
RelatedMemo: &v1pb.MemoRelation_Memo{
|
||||
Name: fmt.Sprintf("%s%s", MemoNamePrefix, relatedMemo.UID),
|
||||
Snippet: relatedMemoSnippet,
|
||||
},
|
||||
Type: convertMemoRelationTypeFromStore(memoRelation.Type),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func convertMemoRelationTypeFromStore(relationType store.MemoRelationType) v1pb.MemoRelation_Type {
|
||||
switch relationType {
|
||||
case store.MemoRelationReference:
|
||||
return v1pb.MemoRelation_REFERENCE
|
||||
case store.MemoRelationComment:
|
||||
return v1pb.MemoRelation_COMMENT
|
||||
default:
|
||||
return v1pb.MemoRelation_TYPE_UNSPECIFIED
|
||||
}
|
||||
}
|
||||
|
||||
func convertMemoRelationTypeToStore(relationType v1pb.MemoRelation_Type) store.MemoRelationType {
|
||||
switch relationType {
|
||||
case v1pb.MemoRelation_COMMENT:
|
||||
return store.MemoRelationComment
|
||||
default:
|
||||
return store.MemoRelationReference
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,372 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
stderrors "errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
"google.golang.org/protobuf/types/known/timestamppb"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
var (
|
||||
errMemoCreatorNotFound = stderrors.New("memo creator not found")
|
||||
errReactionCreatorNotFound = stderrors.New("reaction creator not found")
|
||||
)
|
||||
|
||||
func (s *APIV1Service) convertMemoFromStore(ctx context.Context, memo *store.Memo, reactions []*store.Reaction, attachments []*store.Attachment, relations []*v1pb.MemoRelation) (*v1pb.Memo, error) {
|
||||
creatorMap, err := s.listUsersByID(ctx, []int32{memo.CreatorID})
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to list memo creators")
|
||||
}
|
||||
return s.convertMemoFromStoreWithCreators(ctx, memo, reactions, attachments, relations, creatorMap)
|
||||
}
|
||||
|
||||
func (s *APIV1Service) convertMemoFromStoreWithCreators(ctx context.Context, memo *store.Memo, reactions []*store.Reaction, attachments []*store.Attachment, relations []*v1pb.MemoRelation, creatorMap map[int32]*store.User) (*v1pb.Memo, error) {
|
||||
name := fmt.Sprintf("%s%s", MemoNamePrefix, memo.UID)
|
||||
creator := creatorMap[memo.CreatorID]
|
||||
if creator == nil {
|
||||
return nil, errMemoCreatorNotFound
|
||||
}
|
||||
memoMessage := &v1pb.Memo{
|
||||
Name: name,
|
||||
State: convertStateFromStore(memo.RowStatus),
|
||||
Creator: BuildUserName(creator.Username),
|
||||
CreateTime: timestamppb.New(time.Unix(memo.CreatedTs, 0)),
|
||||
UpdateTime: timestamppb.New(time.Unix(memo.UpdatedTs, 0)),
|
||||
Content: memo.Content,
|
||||
Visibility: convertVisibilityFromStore(memo.Visibility),
|
||||
Pinned: memo.Pinned,
|
||||
}
|
||||
if memo.Payload != nil {
|
||||
memoMessage.Tags = memo.Payload.Tags
|
||||
memoMessage.Property = convertMemoPropertyFromStore(memo.Payload.Property)
|
||||
memoMessage.Location = convertLocationFromStore(memo.Payload.Location)
|
||||
}
|
||||
|
||||
if memo.ParentUID != nil {
|
||||
parentName := fmt.Sprintf("%s%s", MemoNamePrefix, *memo.ParentUID)
|
||||
memoMessage.Parent = &parentName
|
||||
}
|
||||
|
||||
reactionMessages, err := s.convertReactionsFromStoreWithCreators(ctx, reactions, creatorMap)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to convert reactions")
|
||||
}
|
||||
memoMessage.Reactions = reactionMessages
|
||||
|
||||
if relations != nil {
|
||||
memoMessage.Relations = relations
|
||||
} else {
|
||||
memoMessage.Relations = []*v1pb.MemoRelation{}
|
||||
}
|
||||
|
||||
memoMessage.Attachments = []*v1pb.Attachment{}
|
||||
for _, attachment := range attachments {
|
||||
attachmentResponse := convertAttachmentFromStore(attachment)
|
||||
memoMessage.Attachments = append(memoMessage.Attachments, attachmentResponse)
|
||||
}
|
||||
|
||||
snippet, err := s.getMemoContentSnippet(memo.Content)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to get memo content snippet")
|
||||
}
|
||||
memoMessage.Snippet = snippet
|
||||
|
||||
return memoMessage, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) listUsersByIDWithExisting(ctx context.Context, userIDs []int32, existing map[int32]*store.User) (map[int32]*store.User, error) {
|
||||
usersByID := make(map[int32]*store.User, len(existing)+len(userIDs))
|
||||
for userID, user := range existing {
|
||||
if user != nil {
|
||||
usersByID[userID] = user
|
||||
}
|
||||
}
|
||||
|
||||
missingUserIDs := make([]int32, 0, len(userIDs))
|
||||
seenMissingUserIDs := make(map[int32]struct{}, len(userIDs))
|
||||
for _, userID := range userIDs {
|
||||
if _, ok := usersByID[userID]; ok {
|
||||
continue
|
||||
}
|
||||
if _, ok := seenMissingUserIDs[userID]; ok {
|
||||
continue
|
||||
}
|
||||
seenMissingUserIDs[userID] = struct{}{}
|
||||
missingUserIDs = append(missingUserIDs, userID)
|
||||
}
|
||||
|
||||
if len(missingUserIDs) == 0 {
|
||||
return usersByID, nil
|
||||
}
|
||||
|
||||
missingUsersByID, err := s.listUsersByID(ctx, missingUserIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for userID, user := range missingUsersByID {
|
||||
if user != nil {
|
||||
usersByID[userID] = user
|
||||
}
|
||||
}
|
||||
return usersByID, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) convertReactionsFromStoreWithCreators(ctx context.Context, reactions []*store.Reaction, creatorMap map[int32]*store.User) ([]*v1pb.Reaction, error) {
|
||||
if len(reactions) == 0 {
|
||||
return []*v1pb.Reaction{}, nil
|
||||
}
|
||||
|
||||
creatorIDs := make([]int32, 0, len(reactions))
|
||||
for _, reaction := range reactions {
|
||||
creatorIDs = append(creatorIDs, reaction.CreatorID)
|
||||
}
|
||||
creatorsByID, err := s.listUsersByIDWithExisting(ctx, creatorIDs, creatorMap)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
reactionMessages := make([]*v1pb.Reaction, 0, len(reactions))
|
||||
for _, reaction := range reactions {
|
||||
reactionMessage, err := convertReactionFromStoreWithCreators(reaction, creatorsByID)
|
||||
if err != nil {
|
||||
if stderrors.Is(err, errReactionCreatorNotFound) {
|
||||
slog.Warn("Skipping reaction with missing creator",
|
||||
slog.Int64("reaction_id", int64(reaction.ID)),
|
||||
slog.Int64("creator_id", int64(reaction.CreatorID)),
|
||||
slog.String("content_id", reaction.ContentID),
|
||||
)
|
||||
continue
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
reactionMessages = append(reactionMessages, reactionMessage)
|
||||
}
|
||||
return reactionMessages, nil
|
||||
}
|
||||
|
||||
func convertReactionFromStoreWithCreators(reaction *store.Reaction, creatorsByID map[int32]*store.User) (*v1pb.Reaction, error) {
|
||||
creator := creatorsByID[reaction.CreatorID]
|
||||
if creator == nil {
|
||||
return nil, errReactionCreatorNotFound
|
||||
}
|
||||
|
||||
reactionUID := fmt.Sprintf("%d", reaction.ID)
|
||||
return &v1pb.Reaction{
|
||||
Name: fmt.Sprintf("%s/%s%s", reaction.ContentID, ReactionNamePrefix, reactionUID),
|
||||
Creator: BuildUserName(creator.Username),
|
||||
ContentId: reaction.ContentID,
|
||||
ReactionType: reaction.ReactionType,
|
||||
CreateTime: timestamppb.New(time.Unix(reaction.CreatedTs, 0)),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// batchConvertMemoRelations batch-loads relations for a list of memos and returns
|
||||
// a map from memo ID to its converted relations. This avoids N+1 queries when listing memos.
|
||||
func (s *APIV1Service) batchConvertMemoRelations(ctx context.Context, memos []*store.Memo, includeSnippets bool) (map[int32][]*v1pb.MemoRelation, error) {
|
||||
if len(memos) == 0 {
|
||||
return map[int32][]*v1pb.MemoRelation{}, nil
|
||||
}
|
||||
|
||||
currentUser, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to get user")
|
||||
}
|
||||
var memoFilter string
|
||||
if currentUser == nil {
|
||||
memoFilter = `visibility == "PUBLIC"`
|
||||
} else {
|
||||
memoFilter = fmt.Sprintf(`creator_id == %d || visibility in ["PUBLIC", "PROTECTED"]`, currentUser.ID)
|
||||
}
|
||||
|
||||
memoIDs := make([]int32, len(memos))
|
||||
memoIDSet := make(map[int32]bool, len(memos))
|
||||
for i, m := range memos {
|
||||
memoIDs[i] = m.ID
|
||||
memoIDSet[m.ID] = true
|
||||
}
|
||||
|
||||
outgoingRelations, err := s.Store.ListMemoRelations(ctx, &store.FindMemoRelation{
|
||||
SourceMemoIDList: memoIDs,
|
||||
MemoFilter: &memoFilter,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to batch list outgoing memo relations")
|
||||
}
|
||||
incomingRelations, err := s.Store.ListMemoRelations(ctx, &store.FindMemoRelation{
|
||||
RelatedMemoIDList: memoIDs,
|
||||
MemoFilter: &memoFilter,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to batch list incoming memo relations")
|
||||
}
|
||||
allRelations := mergeMemoRelations(outgoingRelations, incomingRelations)
|
||||
|
||||
// Collect all memo IDs referenced in relations that we need to resolve.
|
||||
neededIDs := make(map[int32]bool)
|
||||
for _, r := range allRelations {
|
||||
neededIDs[r.MemoID] = true
|
||||
neededIDs[r.RelatedMemoID] = true
|
||||
}
|
||||
|
||||
// Build ID→UID map from the memos we already have.
|
||||
memoIDToUID := make(map[int32]string, len(memos))
|
||||
memoIDToSnippet := make(map[int32]string, len(memos))
|
||||
for _, m := range memos {
|
||||
memoIDToUID[m.ID] = m.UID
|
||||
if includeSnippets {
|
||||
snippet, err := s.getMemoContentSnippet(m.Content)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to get memo content snippet")
|
||||
}
|
||||
memoIDToSnippet[m.ID] = snippet
|
||||
}
|
||||
delete(neededIDs, m.ID)
|
||||
}
|
||||
|
||||
// Batch fetch any additional memos referenced by relations that we don't already have.
|
||||
if len(neededIDs) > 0 {
|
||||
extraIDs := make([]int32, 0, len(neededIDs))
|
||||
for id := range neededIDs {
|
||||
extraIDs = append(extraIDs, id)
|
||||
}
|
||||
extraFind := &store.FindMemo{IDList: extraIDs, ExcludeContent: !includeSnippets}
|
||||
extraMemos, err := s.Store.ListMemos(ctx, extraFind)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to batch fetch related memos")
|
||||
}
|
||||
for _, m := range extraMemos {
|
||||
memoIDToUID[m.ID] = m.UID
|
||||
if includeSnippets {
|
||||
snippet, err := s.getMemoContentSnippet(m.Content)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to get related memo content snippet")
|
||||
}
|
||||
memoIDToSnippet[m.ID] = snippet
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Build the result map: memo ID → its relations (both directions).
|
||||
result := make(map[int32][]*v1pb.MemoRelation, len(memos))
|
||||
for _, r := range allRelations {
|
||||
memoUID, ok1 := memoIDToUID[r.MemoID]
|
||||
relatedUID, ok2 := memoIDToUID[r.RelatedMemoID]
|
||||
if !ok1 || !ok2 {
|
||||
continue
|
||||
}
|
||||
|
||||
relation := &v1pb.MemoRelation{
|
||||
Memo: &v1pb.MemoRelation_Memo{
|
||||
Name: fmt.Sprintf("%s%s", MemoNamePrefix, memoUID),
|
||||
Snippet: memoIDToSnippet[r.MemoID],
|
||||
},
|
||||
RelatedMemo: &v1pb.MemoRelation_Memo{
|
||||
Name: fmt.Sprintf("%s%s", MemoNamePrefix, relatedUID),
|
||||
Snippet: memoIDToSnippet[r.RelatedMemoID],
|
||||
},
|
||||
Type: convertMemoRelationTypeFromStore(r.Type),
|
||||
}
|
||||
|
||||
// Add to the memo that owns this relation (both directions).
|
||||
if memoIDSet[r.MemoID] {
|
||||
result[r.MemoID] = append(result[r.MemoID], relation)
|
||||
}
|
||||
if memoIDSet[r.RelatedMemoID] {
|
||||
result[r.RelatedMemoID] = append(result[r.RelatedMemoID], relation)
|
||||
}
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// loadMemoRelations loads relations for a single memo and converts them to API format.
|
||||
func (s *APIV1Service) loadMemoRelations(ctx context.Context, memo *store.Memo) ([]*v1pb.MemoRelation, error) {
|
||||
relationMap, err := s.batchConvertMemoRelations(ctx, []*store.Memo{memo}, true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return relationMap[memo.ID], nil
|
||||
}
|
||||
|
||||
func mergeMemoRelations(groups ...[]*store.MemoRelation) []*store.MemoRelation {
|
||||
seen := make(map[string]struct{})
|
||||
merged := make([]*store.MemoRelation, 0)
|
||||
for _, relations := range groups {
|
||||
for _, relation := range relations {
|
||||
key := fmt.Sprintf("%d:%d:%s", relation.MemoID, relation.RelatedMemoID, relation.Type)
|
||||
if _, ok := seen[key]; ok {
|
||||
continue
|
||||
}
|
||||
seen[key] = struct{}{}
|
||||
merged = append(merged, relation)
|
||||
}
|
||||
}
|
||||
return merged
|
||||
}
|
||||
|
||||
func convertMemoPropertyFromStore(property *storepb.MemoPayload_Property) *v1pb.Memo_Property {
|
||||
if property == nil {
|
||||
return nil
|
||||
}
|
||||
return &v1pb.Memo_Property{
|
||||
HasLink: property.HasLink,
|
||||
HasTaskList: property.HasTaskList,
|
||||
HasCode: property.HasCode,
|
||||
HasIncompleteTasks: property.HasIncompleteTasks,
|
||||
Title: property.Title,
|
||||
}
|
||||
}
|
||||
|
||||
func convertLocationFromStore(location *storepb.MemoPayload_Location) *v1pb.Location {
|
||||
if location == nil {
|
||||
return nil
|
||||
}
|
||||
return &v1pb.Location{
|
||||
Placeholder: location.Placeholder,
|
||||
Latitude: location.Latitude,
|
||||
Longitude: location.Longitude,
|
||||
}
|
||||
}
|
||||
|
||||
func convertLocationToStore(location *v1pb.Location) *storepb.MemoPayload_Location {
|
||||
if location == nil {
|
||||
return nil
|
||||
}
|
||||
return &storepb.MemoPayload_Location{
|
||||
Placeholder: location.Placeholder,
|
||||
Latitude: location.Latitude,
|
||||
Longitude: location.Longitude,
|
||||
}
|
||||
}
|
||||
|
||||
func convertVisibilityFromStore(visibility store.Visibility) v1pb.Visibility {
|
||||
switch visibility {
|
||||
case store.Private:
|
||||
return v1pb.Visibility_PRIVATE
|
||||
case store.Protected:
|
||||
return v1pb.Visibility_PROTECTED
|
||||
case store.Public:
|
||||
return v1pb.Visibility_PUBLIC
|
||||
default:
|
||||
return v1pb.Visibility_VISIBILITY_UNSPECIFIED
|
||||
}
|
||||
}
|
||||
|
||||
func convertVisibilityToStore(visibility v1pb.Visibility) store.Visibility {
|
||||
switch visibility {
|
||||
case v1pb.Visibility_PROTECTED:
|
||||
return store.Protected
|
||||
case v1pb.Visibility_PUBLIC:
|
||||
return store.Public
|
||||
default:
|
||||
return store.Private
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1 @@
|
||||
package v1
|
||||
@@ -0,0 +1,226 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
stderrors "errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
"google.golang.org/protobuf/types/known/emptypb"
|
||||
"google.golang.org/protobuf/types/known/timestamppb"
|
||||
|
||||
"github.com/lithammer/shortuuid/v4"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
// CreateMemoShare creates an opaque share link for a memo.
|
||||
// Only the memo's creator or an admin may call this.
|
||||
func (s *APIV1Service) CreateMemoShare(ctx context.Context, request *v1pb.CreateMemoShareRequest) (*v1pb.MemoShare, error) {
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get user")
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
|
||||
memoUID, err := ExtractMemoUIDFromName(request.Parent)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid memo name: %v", err)
|
||||
}
|
||||
memo, err := s.Store.GetMemo(ctx, &store.FindMemo{UID: &memoUID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get memo")
|
||||
}
|
||||
if memo == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "memo not found")
|
||||
}
|
||||
if memo.CreatorID != user.ID && !isSuperUser(user) {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
|
||||
var expiresTs *int64
|
||||
if request.MemoShare != nil && request.MemoShare.ExpireTime != nil {
|
||||
ts := request.MemoShare.ExpireTime.AsTime().Unix()
|
||||
if ts <= time.Now().Unix() {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "expire_time must be in the future")
|
||||
}
|
||||
expiresTs = &ts
|
||||
}
|
||||
|
||||
// Generate a URL-safe token using shortuuid (base57-encoded UUID v4, 22 chars, 122-bit entropy).
|
||||
ms, err := s.Store.CreateMemoShare(ctx, &store.MemoShare{
|
||||
UID: shortuuid.New(),
|
||||
MemoID: memo.ID,
|
||||
CreatorID: user.ID,
|
||||
ExpiresTs: expiresTs,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to create memo share")
|
||||
}
|
||||
|
||||
return convertMemoShareFromStore(ms, memo.UID), nil
|
||||
}
|
||||
|
||||
// ListMemoShares lists all share links for a memo.
|
||||
// Only the memo's creator or an admin may call this.
|
||||
func (s *APIV1Service) ListMemoShares(ctx context.Context, request *v1pb.ListMemoSharesRequest) (*v1pb.ListMemoSharesResponse, error) {
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get user")
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
|
||||
memoUID, err := ExtractMemoUIDFromName(request.Parent)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid memo name: %v", err)
|
||||
}
|
||||
memo, err := s.Store.GetMemo(ctx, &store.FindMemo{UID: &memoUID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get memo")
|
||||
}
|
||||
if memo == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "memo not found")
|
||||
}
|
||||
if memo.CreatorID != user.ID && !isSuperUser(user) {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
|
||||
shares, err := s.Store.ListMemoShares(ctx, &store.FindMemoShare{MemoID: &memo.ID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to list memo shares")
|
||||
}
|
||||
|
||||
response := &v1pb.ListMemoSharesResponse{}
|
||||
for _, ms := range shares {
|
||||
response.MemoShares = append(response.MemoShares, convertMemoShareFromStore(ms, memo.UID))
|
||||
}
|
||||
return response, nil
|
||||
}
|
||||
|
||||
// DeleteMemoShare revokes a share link.
|
||||
// Only the memo's creator or an admin may call this.
|
||||
func (s *APIV1Service) DeleteMemoShare(ctx context.Context, request *v1pb.DeleteMemoShareRequest) (*emptypb.Empty, error) {
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get user")
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
|
||||
// name format: memos/{memoUID}/shares/{shareToken}
|
||||
tokens, err := GetNameParentTokens(request.Name, MemoNamePrefix, MemoShareNamePrefix)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid share name: %v", err)
|
||||
}
|
||||
memoUID, shareToken := tokens[0], tokens[1]
|
||||
|
||||
memo, err := s.Store.GetMemo(ctx, &store.FindMemo{UID: &memoUID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get memo")
|
||||
}
|
||||
if memo == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "memo not found")
|
||||
}
|
||||
if memo.CreatorID != user.ID && !isSuperUser(user) {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
|
||||
ms, err := s.Store.GetMemoShare(ctx, &store.FindMemoShare{UID: &shareToken})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get memo share")
|
||||
}
|
||||
if ms == nil || ms.MemoID != memo.ID {
|
||||
return nil, status.Errorf(codes.NotFound, "memo share not found")
|
||||
}
|
||||
|
||||
if err := s.Store.DeleteMemoShare(ctx, &store.DeleteMemoShare{UID: &shareToken}); err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to delete memo share")
|
||||
}
|
||||
return &emptypb.Empty{}, nil
|
||||
}
|
||||
|
||||
// GetMemoByShare resolves a share token to its memo. No authentication required.
|
||||
// Returns NOT_FOUND for invalid or expired tokens (no information leakage).
|
||||
func (s *APIV1Service) GetMemoByShare(ctx context.Context, request *v1pb.GetMemoByShareRequest) (*v1pb.Memo, error) {
|
||||
ms, err := s.getActiveMemoShare(ctx, request.ShareId)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
memo, err := s.Store.GetMemo(ctx, &store.FindMemo{ID: &ms.MemoID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get memo")
|
||||
}
|
||||
// Treat archived or missing memos the same as an invalid token — no information leakage.
|
||||
if memo == nil || memo.RowStatus == store.Archived {
|
||||
return nil, status.Errorf(codes.NotFound, "not found")
|
||||
}
|
||||
|
||||
reactions, err := s.Store.ListReactions(ctx, &store.FindReaction{
|
||||
ContentID: stringPointer(fmt.Sprintf("%s%s", MemoNamePrefix, memo.UID)),
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to list reactions")
|
||||
}
|
||||
|
||||
attachments, err := s.Store.ListAttachments(ctx, &store.FindAttachment{MemoID: &memo.ID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to list attachments")
|
||||
}
|
||||
relations, err := s.batchConvertMemoRelations(ctx, []*store.Memo{memo}, true)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to load memo relations")
|
||||
}
|
||||
|
||||
memoMessage, err := s.convertMemoFromStore(ctx, memo, reactions, attachments, relations[memo.ID])
|
||||
if err != nil {
|
||||
if stderrors.Is(err, errMemoCreatorNotFound) {
|
||||
return nil, status.Errorf(codes.NotFound, "not found")
|
||||
}
|
||||
return nil, errors.Wrap(err, "failed to convert memo")
|
||||
}
|
||||
return memoMessage, nil
|
||||
}
|
||||
|
||||
// isMemoShareExpired returns true if the share has a defined expiry that has already passed.
|
||||
func isMemoShareExpired(ms *store.MemoShare) bool {
|
||||
return ms.ExpiresTs != nil && time.Now().Unix() > *ms.ExpiresTs
|
||||
}
|
||||
|
||||
func (s *APIV1Service) getActiveMemoShare(ctx context.Context, shareID string) (*store.MemoShare, error) {
|
||||
ms, err := s.Store.GetMemoShare(ctx, &store.FindMemoShare{UID: &shareID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get memo share")
|
||||
}
|
||||
if ms == nil || isMemoShareExpired(ms) {
|
||||
return nil, status.Errorf(codes.NotFound, "not found")
|
||||
}
|
||||
return ms, nil
|
||||
}
|
||||
|
||||
func stringPointer(s string) *string {
|
||||
return &s
|
||||
}
|
||||
|
||||
// convertMemoShareFromStore converts a store MemoShare to the proto MemoShare message.
|
||||
// name format: memos/{memoUID}/shares/{shareToken}.
|
||||
func convertMemoShareFromStore(ms *store.MemoShare, memoUID string) *v1pb.MemoShare {
|
||||
name := fmt.Sprintf("%s%s/%s%s", MemoNamePrefix, memoUID, MemoShareNamePrefix, ms.UID)
|
||||
pb := &v1pb.MemoShare{
|
||||
Name: name,
|
||||
CreateTime: timestamppb.New(time.Unix(ms.CreatedTs, 0)),
|
||||
}
|
||||
if ms.ExpiresTs != nil {
|
||||
pb.ExpireTime = timestamppb.New(time.Unix(*ms.ExpiresTs, 0))
|
||||
}
|
||||
return pb
|
||||
}
|
||||
@@ -0,0 +1,78 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func (s *APIV1Service) touchMemoUpdatedTimestamp(ctx context.Context, memoID int32) error {
|
||||
updatedTs := time.Now().Unix()
|
||||
if err := s.Store.UpdateMemo(ctx, &store.UpdateMemo{
|
||||
ID: memoID,
|
||||
UpdatedTs: &updatedTs,
|
||||
}); err != nil {
|
||||
return status.Errorf(codes.Internal, "failed to update memo timestamp")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) buildUpdatedMemoState(ctx context.Context, memoID int32) (*store.Memo, *store.Memo, *v1pb.Memo, error) {
|
||||
memo, err := s.Store.GetMemo(ctx, &store.FindMemo{ID: &memoID})
|
||||
if err != nil {
|
||||
return nil, nil, nil, errors.Wrap(err, "failed to get memo")
|
||||
}
|
||||
if memo == nil {
|
||||
return nil, nil, nil, errors.New("memo not found")
|
||||
}
|
||||
|
||||
memoName := buildMemoName(memo.UID)
|
||||
reactions, err := s.Store.ListReactions(ctx, &store.FindReaction{
|
||||
ContentID: &memoName,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, nil, nil, errors.Wrap(err, "failed to list reactions")
|
||||
}
|
||||
attachments, err := s.Store.ListAttachments(ctx, &store.FindAttachment{
|
||||
MemoID: &memo.ID,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, nil, nil, errors.Wrap(err, "failed to list attachments")
|
||||
}
|
||||
relations, err := s.loadMemoRelations(ctx, memo)
|
||||
if err != nil {
|
||||
return nil, nil, nil, errors.Wrap(err, "failed to load memo relations")
|
||||
}
|
||||
memoMessage, err := s.convertMemoFromStore(ctx, memo, reactions, attachments, relations)
|
||||
if err != nil {
|
||||
return nil, nil, nil, errors.Wrap(err, "failed to convert memo")
|
||||
}
|
||||
|
||||
var parentMemo *store.Memo
|
||||
if memo.ParentUID != nil {
|
||||
parentMemo, _ = s.Store.GetMemo(ctx, &store.FindMemo{UID: memo.ParentUID})
|
||||
}
|
||||
|
||||
return memo, parentMemo, memoMessage, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) dispatchMemoUpdatedSideEffects(ctx context.Context, memo *store.Memo, parentMemo *store.Memo, memoMessage *v1pb.Memo) {
|
||||
if err := s.DispatchMemoUpdatedWebhook(ctx, memoMessage); err != nil {
|
||||
slog.Warn("Failed to dispatch memo updated webhook", slog.Any("err", err))
|
||||
}
|
||||
|
||||
s.SSEHub.Broadcast(&SSEEvent{
|
||||
Type: SSEEventMemoUpdated,
|
||||
Name: memoMessage.Name,
|
||||
Parent: memoMessage.GetParent(),
|
||||
Visibility: memo.Visibility,
|
||||
CreatorID: resolveSSECreatorID(memo, parentMemo),
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
|
||||
"github.com/usememos/memos/server/notification"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func (s *APIV1Service) createInboxWithEmailNotification(ctx context.Context, inbox *store.Inbox) (*store.Inbox, error) {
|
||||
createdInbox, err := s.Store.CreateInbox(ctx, inbox)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
s.dispatchInboxEmailNotificationBestEffort(ctx, createdInbox)
|
||||
return createdInbox, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) dispatchInboxEmailNotificationBestEffort(ctx context.Context, inbox *store.Inbox) {
|
||||
dispatcher := notification.NewEmailDispatcher(s.Profile, s.Store, s.NotificationEmailSender)
|
||||
if err := dispatcher.DispatchInboxEmail(ctx, inbox); err != nil {
|
||||
slog.Warn("Failed to dispatch inbox email notification",
|
||||
slog.Any("err", err),
|
||||
slog.Int64("inbox_id", int64(inbox.ID)),
|
||||
slog.Int64("receiver_id", int64(inbox.ReceiverID)))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
"google.golang.org/protobuf/types/known/emptypb"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func (s *APIV1Service) ListMemoReactions(ctx context.Context, request *v1pb.ListMemoReactionsRequest) (*v1pb.ListMemoReactionsResponse, error) {
|
||||
// Extract memo UID and check visibility.
|
||||
memoUID, err := ExtractMemoUIDFromName(request.Name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid memo name: %v", err)
|
||||
}
|
||||
memo, err := s.Store.GetMemo(ctx, &store.FindMemo{UID: &memoUID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get memo: %v", err)
|
||||
}
|
||||
if memo == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "memo not found")
|
||||
}
|
||||
|
||||
// Check memo visibility.
|
||||
if memo.Visibility != store.Public {
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get user")
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
if memo.Visibility == store.Private && memo.CreatorID != user.ID && !isSuperUser(user) {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
}
|
||||
|
||||
reactions, err := s.Store.ListReactions(ctx, &store.FindReaction{
|
||||
ContentID: &request.Name,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to list reactions")
|
||||
}
|
||||
|
||||
response := &v1pb.ListMemoReactionsResponse{
|
||||
Reactions: []*v1pb.Reaction{},
|
||||
}
|
||||
response.Reactions, err = s.convertReactionsFromStoreWithCreators(ctx, reactions, nil)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to convert reactions")
|
||||
}
|
||||
return response, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) UpsertMemoReaction(ctx context.Context, request *v1pb.UpsertMemoReactionRequest) (*v1pb.Reaction, error) {
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user")
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
|
||||
// Extract memo UID and check visibility before allowing reaction.
|
||||
memoUID, err := ExtractMemoUIDFromName(request.Reaction.ContentId)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid memo name: %v", err)
|
||||
}
|
||||
memo, err := s.Store.GetMemo(ctx, &store.FindMemo{UID: &memoUID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get memo: %v", err)
|
||||
}
|
||||
if memo == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "memo not found")
|
||||
}
|
||||
|
||||
// Check memo visibility.
|
||||
if memo.Visibility == store.Private && memo.CreatorID != user.ID && !isSuperUser(user) {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
|
||||
reaction, err := s.Store.UpsertReaction(ctx, &store.Reaction{
|
||||
CreatorID: user.ID,
|
||||
ContentID: request.Reaction.ContentId,
|
||||
ReactionType: request.Reaction.ReactionType,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to upsert reaction")
|
||||
}
|
||||
|
||||
reactionMessage, err := s.convertReactionFromStore(ctx, reaction)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to convert reaction")
|
||||
}
|
||||
|
||||
// Broadcast live refresh event (reaction belongs to a memo).
|
||||
var parentMemo *store.Memo
|
||||
if memo.ParentUID != nil {
|
||||
parentMemo, _ = s.Store.GetMemo(ctx, &store.FindMemo{UID: memo.ParentUID})
|
||||
}
|
||||
s.SSEHub.Broadcast(buildMemoReactionSSEEvent(SSEEventReactionUpserted, request.Reaction.ContentId, memo, parentMemo))
|
||||
|
||||
return reactionMessage, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) DeleteMemoReaction(ctx context.Context, request *v1pb.DeleteMemoReactionRequest) (*emptypb.Empty, error) {
|
||||
user, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.Unauthenticated, "user not authenticated")
|
||||
}
|
||||
|
||||
_, reactionID, err := ExtractMemoReactionIDFromName(request.Name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid reaction name: %v", err)
|
||||
}
|
||||
|
||||
// Get reaction and check ownership.
|
||||
reaction, err := s.Store.GetReaction(ctx, &store.FindReaction{
|
||||
ID: &reactionID,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get reaction")
|
||||
}
|
||||
if reaction == nil {
|
||||
// Return permission denied to avoid revealing if reaction exists.
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
|
||||
if reaction.CreatorID != user.ID && !isSuperUser(user) {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
|
||||
if err := s.Store.DeleteReaction(ctx, &store.DeleteReaction{
|
||||
ID: reactionID,
|
||||
}); err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to delete reaction")
|
||||
}
|
||||
|
||||
memoUID, err := ExtractMemoUIDFromName(reaction.ContentID)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid memo name: %v", err)
|
||||
}
|
||||
memo, err := s.Store.GetMemo(ctx, &store.FindMemo{UID: &memoUID})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get memo")
|
||||
}
|
||||
|
||||
// Broadcast live refresh event (reaction belongs to a memo).
|
||||
var parentMemo *store.Memo
|
||||
if memo != nil && memo.ParentUID != nil {
|
||||
parentMemo, _ = s.Store.GetMemo(ctx, &store.FindMemo{UID: memo.ParentUID})
|
||||
}
|
||||
s.SSEHub.Broadcast(buildMemoReactionSSEEvent(SSEEventReactionDeleted, reaction.ContentID, memo, parentMemo))
|
||||
|
||||
return &emptypb.Empty{}, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) convertReactionFromStore(ctx context.Context, reaction *store.Reaction) (*v1pb.Reaction, error) {
|
||||
creatorsByID, err := s.listUsersByIDWithExisting(ctx, []int32{reaction.CreatorID}, nil)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get reaction creator")
|
||||
}
|
||||
reactionMessage, err := convertReactionFromStoreWithCreators(reaction, creatorsByID)
|
||||
if err != nil {
|
||||
slog.Warn("Failed to convert reaction with missing creator",
|
||||
slog.Int64("reaction_id", int64(reaction.ID)),
|
||||
slog.Int64("creator_id", int64(reaction.CreatorID)),
|
||||
slog.String("content_id", reaction.ContentID),
|
||||
)
|
||||
return nil, status.Errorf(codes.NotFound, "reaction creator not found")
|
||||
}
|
||||
return reactionMessage, nil
|
||||
}
|
||||
@@ -0,0 +1,149 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/lithammer/shortuuid/v4"
|
||||
"github.com/pkg/errors"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
"github.com/usememos/memos/internal/base"
|
||||
"github.com/usememos/memos/internal/util"
|
||||
)
|
||||
|
||||
const (
|
||||
InstanceSettingNamePrefix = "instance/settings/"
|
||||
UserNamePrefix = "users/"
|
||||
MemoNamePrefix = "memos/"
|
||||
MemoShareNamePrefix = "shares/"
|
||||
AttachmentNamePrefix = "attachments/"
|
||||
ReactionNamePrefix = "reactions/"
|
||||
InboxNamePrefix = "inboxes/"
|
||||
IdentityProviderNamePrefix = "identity-providers/"
|
||||
WebhookNamePrefix = "webhooks/"
|
||||
)
|
||||
|
||||
// GetNameParentTokens returns the tokens from a resource name.
|
||||
func GetNameParentTokens(name string, tokenPrefixes ...string) ([]string, error) {
|
||||
parts := strings.Split(name, "/")
|
||||
if len(parts) != 2*len(tokenPrefixes) {
|
||||
return nil, errors.Errorf("invalid request %q", name)
|
||||
}
|
||||
|
||||
var tokens []string
|
||||
for i, tokenPrefix := range tokenPrefixes {
|
||||
if fmt.Sprintf("%s/", parts[2*i]) != tokenPrefix {
|
||||
return nil, errors.Errorf("invalid prefix %q in request %q", tokenPrefix, name)
|
||||
}
|
||||
if parts[2*i+1] == "" {
|
||||
return nil, errors.Errorf("invalid request %q with empty prefix %q", name, tokenPrefix)
|
||||
}
|
||||
tokens = append(tokens, parts[2*i+1])
|
||||
}
|
||||
return tokens, nil
|
||||
}
|
||||
|
||||
func ExtractInstanceSettingKeyFromName(name string) (string, error) {
|
||||
const prefix = "instance/settings/"
|
||||
if !strings.HasPrefix(name, prefix) {
|
||||
return "", errors.Errorf("invalid instance setting name: expected prefix %q, got %q", prefix, name)
|
||||
}
|
||||
|
||||
settingKey := strings.TrimPrefix(name, prefix)
|
||||
if settingKey == "" {
|
||||
return "", errors.Errorf("invalid instance setting name: empty setting key in %q", name)
|
||||
}
|
||||
|
||||
// Ensure there are no additional path segments
|
||||
if strings.Contains(settingKey, "/") {
|
||||
return "", errors.Errorf("invalid instance setting name: setting key cannot contain '/' in %q", name)
|
||||
}
|
||||
|
||||
return settingKey, nil
|
||||
}
|
||||
|
||||
// ExtractUserIDFromName returns the uid from a resource name.
|
||||
func ExtractUserIDFromName(name string) (int32, error) {
|
||||
tokens, err := GetNameParentTokens(name, UserNamePrefix)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
id, err := util.ConvertStringToInt32(tokens[0])
|
||||
if err != nil {
|
||||
return 0, errors.Errorf("invalid user ID %q", tokens[0])
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// ExtractMemoUIDFromName returns the memo UID from a resource name.
|
||||
// e.g., "memos/uuid" -> "uuid".
|
||||
func ExtractMemoUIDFromName(name string) (string, error) {
|
||||
tokens, err := GetNameParentTokens(name, MemoNamePrefix)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
id := tokens[0]
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// ExtractAttachmentUIDFromName returns the attachment UID from a resource name.
|
||||
func ExtractAttachmentUIDFromName(name string) (string, error) {
|
||||
tokens, err := GetNameParentTokens(name, AttachmentNamePrefix)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
id := tokens[0]
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// ExtractMemoReactionIDFromName returns the memo UID and reaction ID from a resource name.
|
||||
// e.g., "memos/abc/reactions/123" -> ("abc", 123).
|
||||
func ExtractMemoReactionIDFromName(name string) (string, int32, error) {
|
||||
tokens, err := GetNameParentTokens(name, MemoNamePrefix, ReactionNamePrefix)
|
||||
if err != nil {
|
||||
return "", 0, err
|
||||
}
|
||||
memoUID := tokens[0]
|
||||
reactionID, err := util.ConvertStringToInt32(tokens[1])
|
||||
if err != nil {
|
||||
return "", 0, errors.Errorf("invalid reaction ID %q", tokens[1])
|
||||
}
|
||||
return memoUID, reactionID, nil
|
||||
}
|
||||
|
||||
// ExtractInboxIDFromName returns the inbox ID from a resource name.
|
||||
func ExtractInboxIDFromName(name string) (int32, error) {
|
||||
tokens, err := GetNameParentTokens(name, InboxNamePrefix)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
id, err := util.ConvertStringToInt32(tokens[0])
|
||||
if err != nil {
|
||||
return 0, errors.Errorf("invalid inbox ID %q", tokens[0])
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
func ExtractIdentityProviderUIDFromName(name string) (string, error) {
|
||||
tokens, err := GetNameParentTokens(name, IdentityProviderNamePrefix)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return tokens[0], nil
|
||||
}
|
||||
|
||||
// ValidateAndGenerateUID validates a user-provided UID or generates a new one.
|
||||
// If provided is empty, a new shortuuid is generated.
|
||||
// If provided is non-empty, it is validated against base.UIDMatcher.
|
||||
func ValidateAndGenerateUID(provided string) (string, error) {
|
||||
uid := strings.TrimSpace(provided)
|
||||
if uid == "" {
|
||||
return shortuuid.New(), nil
|
||||
}
|
||||
if !base.UIDMatcher.MatchString(uid) {
|
||||
return "", status.Errorf(codes.InvalidArgument, "invalid ID format: must be 1-36 characters, alphanumeric and hyphens only, cannot start or end with hyphen")
|
||||
}
|
||||
return uid, nil
|
||||
}
|
||||
@@ -0,0 +1,360 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
"google.golang.org/protobuf/types/known/emptypb"
|
||||
|
||||
"github.com/usememos/memos/internal/filter"
|
||||
"github.com/usememos/memos/internal/util"
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
// Helper function to extract user and shortcut ID from shortcut resource name.
|
||||
// Format: users/{user}/shortcuts/{shortcut}.
|
||||
func (s *APIV1Service) extractUserAndShortcutIDFromName(ctx context.Context, name string) (*store.User, string, error) {
|
||||
parts := strings.Split(name, "/")
|
||||
if len(parts) != 4 || parts[0] != "users" || parts[2] != "shortcuts" {
|
||||
return nil, "", errors.Errorf("invalid shortcut name format: %s", name)
|
||||
}
|
||||
|
||||
user, err := ResolveUserByName(ctx, s.Store, BuildUserName(parts[1]))
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
if user == nil {
|
||||
return nil, "", errors.Errorf("user not found: %s", parts[1])
|
||||
}
|
||||
|
||||
shortcutID := parts[3]
|
||||
if shortcutID == "" {
|
||||
return nil, "", errors.Errorf("empty shortcut ID in name: %s", name)
|
||||
}
|
||||
|
||||
return user, shortcutID, nil
|
||||
}
|
||||
|
||||
// Helper function to construct shortcut resource name.
|
||||
func constructShortcutName(username string, shortcutID string) string {
|
||||
return fmt.Sprintf("%s/shortcuts/%s", BuildUserName(username), shortcutID)
|
||||
}
|
||||
|
||||
func (s *APIV1Service) ListShortcuts(ctx context.Context, request *v1pb.ListShortcutsRequest) (*v1pb.ListShortcutsResponse, error) {
|
||||
user, err := ResolveUserByName(ctx, s.Store, request.Parent)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid user name: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "user not found")
|
||||
}
|
||||
userID := user.ID
|
||||
|
||||
currentUser, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if currentUser == nil || currentUser.ID != userID {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
|
||||
userSetting, err := s.Store.GetUserSetting(ctx, &store.FindUserSetting{
|
||||
UserID: &userID,
|
||||
Key: storepb.UserSetting_SHORTCUTS,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get user setting: %v", err)
|
||||
}
|
||||
if userSetting == nil {
|
||||
return &v1pb.ListShortcutsResponse{
|
||||
Shortcuts: []*v1pb.Shortcut{},
|
||||
}, nil
|
||||
}
|
||||
|
||||
shortcutsUserSetting := userSetting.GetShortcuts()
|
||||
shortcuts := []*v1pb.Shortcut{}
|
||||
for _, shortcut := range shortcutsUserSetting.GetShortcuts() {
|
||||
shortcuts = append(shortcuts, &v1pb.Shortcut{
|
||||
Name: constructShortcutName(user.Username, shortcut.GetId()),
|
||||
Title: shortcut.GetTitle(),
|
||||
Filter: shortcut.GetFilter(),
|
||||
})
|
||||
}
|
||||
|
||||
return &v1pb.ListShortcutsResponse{
|
||||
Shortcuts: shortcuts,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) GetShortcut(ctx context.Context, request *v1pb.GetShortcutRequest) (*v1pb.Shortcut, error) {
|
||||
user, shortcutID, err := s.extractUserAndShortcutIDFromName(ctx, request.Name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid shortcut name: %v", err)
|
||||
}
|
||||
userID := user.ID
|
||||
|
||||
currentUser, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if currentUser == nil || currentUser.ID != userID {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
|
||||
userSetting, err := s.Store.GetUserSetting(ctx, &store.FindUserSetting{
|
||||
UserID: &userID,
|
||||
Key: storepb.UserSetting_SHORTCUTS,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if userSetting == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "shortcut not found")
|
||||
}
|
||||
|
||||
shortcutsUserSetting := userSetting.GetShortcuts()
|
||||
for _, shortcut := range shortcutsUserSetting.GetShortcuts() {
|
||||
if shortcut.GetId() == shortcutID {
|
||||
return &v1pb.Shortcut{
|
||||
Name: constructShortcutName(user.Username, shortcut.GetId()),
|
||||
Title: shortcut.GetTitle(),
|
||||
Filter: shortcut.GetFilter(),
|
||||
}, nil
|
||||
}
|
||||
}
|
||||
|
||||
return nil, status.Errorf(codes.NotFound, "shortcut not found")
|
||||
}
|
||||
|
||||
func (s *APIV1Service) CreateShortcut(ctx context.Context, request *v1pb.CreateShortcutRequest) (*v1pb.Shortcut, error) {
|
||||
user, err := ResolveUserByName(ctx, s.Store, request.Parent)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid user name: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "user not found")
|
||||
}
|
||||
userID := user.ID
|
||||
|
||||
currentUser, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if currentUser == nil || currentUser.ID != userID {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
|
||||
newShortcut := &storepb.ShortcutsUserSetting_Shortcut{
|
||||
Id: util.GenUUID(),
|
||||
Title: request.Shortcut.GetTitle(),
|
||||
Filter: request.Shortcut.GetFilter(),
|
||||
}
|
||||
if newShortcut.Title == "" {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "title is required")
|
||||
}
|
||||
if err := s.validateFilter(ctx, newShortcut.Filter); err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid filter: %v", err)
|
||||
}
|
||||
if request.ValidateOnly {
|
||||
return &v1pb.Shortcut{
|
||||
Name: constructShortcutName(user.Username, newShortcut.GetId()),
|
||||
Title: newShortcut.GetTitle(),
|
||||
Filter: newShortcut.GetFilter(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
userSetting, err := s.Store.GetUserSetting(ctx, &store.FindUserSetting{
|
||||
UserID: &userID,
|
||||
Key: storepb.UserSetting_SHORTCUTS,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if userSetting == nil {
|
||||
userSetting = &storepb.UserSetting{
|
||||
UserId: userID,
|
||||
Key: storepb.UserSetting_SHORTCUTS,
|
||||
Value: &storepb.UserSetting_Shortcuts{
|
||||
Shortcuts: &storepb.ShortcutsUserSetting{
|
||||
Shortcuts: []*storepb.ShortcutsUserSetting_Shortcut{},
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
shortcutsUserSetting := userSetting.GetShortcuts()
|
||||
shortcuts := shortcutsUserSetting.GetShortcuts()
|
||||
shortcuts = append(shortcuts, newShortcut)
|
||||
shortcutsUserSetting.Shortcuts = shortcuts
|
||||
|
||||
userSetting.Value = &storepb.UserSetting_Shortcuts{
|
||||
Shortcuts: shortcutsUserSetting,
|
||||
}
|
||||
|
||||
_, err = s.Store.UpsertUserSetting(ctx, userSetting)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to upsert user setting: %v", err)
|
||||
}
|
||||
|
||||
return &v1pb.Shortcut{
|
||||
Name: constructShortcutName(user.Username, newShortcut.GetId()),
|
||||
Title: newShortcut.GetTitle(),
|
||||
Filter: newShortcut.GetFilter(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) UpdateShortcut(ctx context.Context, request *v1pb.UpdateShortcutRequest) (*v1pb.Shortcut, error) {
|
||||
user, shortcutID, err := s.extractUserAndShortcutIDFromName(ctx, request.Shortcut.Name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid shortcut name: %v", err)
|
||||
}
|
||||
userID := user.ID
|
||||
|
||||
currentUser, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if currentUser == nil || currentUser.ID != userID {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
if request.UpdateMask == nil || len(request.UpdateMask.Paths) == 0 {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "update mask is required")
|
||||
}
|
||||
|
||||
userSetting, err := s.Store.GetUserSetting(ctx, &store.FindUserSetting{
|
||||
UserID: &userID,
|
||||
Key: storepb.UserSetting_SHORTCUTS,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if userSetting == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "shortcut not found")
|
||||
}
|
||||
|
||||
shortcutsUserSetting := userSetting.GetShortcuts()
|
||||
shortcuts := shortcutsUserSetting.GetShortcuts()
|
||||
var foundShortcut *storepb.ShortcutsUserSetting_Shortcut
|
||||
newShortcuts := make([]*storepb.ShortcutsUserSetting_Shortcut, 0, len(shortcuts))
|
||||
for _, shortcut := range shortcuts {
|
||||
if shortcut.GetId() == shortcutID {
|
||||
foundShortcut = shortcut
|
||||
for _, field := range request.UpdateMask.Paths {
|
||||
if field == "title" {
|
||||
if request.Shortcut.GetTitle() == "" {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "title is required")
|
||||
}
|
||||
shortcut.Title = request.Shortcut.GetTitle()
|
||||
} else if field == "filter" {
|
||||
if err := s.validateFilter(ctx, request.Shortcut.GetFilter()); err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid filter: %v", err)
|
||||
}
|
||||
shortcut.Filter = request.Shortcut.GetFilter()
|
||||
}
|
||||
}
|
||||
}
|
||||
newShortcuts = append(newShortcuts, shortcut)
|
||||
}
|
||||
|
||||
if foundShortcut == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "shortcut not found")
|
||||
}
|
||||
|
||||
shortcutsUserSetting.Shortcuts = newShortcuts
|
||||
userSetting.Value = &storepb.UserSetting_Shortcuts{
|
||||
Shortcuts: shortcutsUserSetting,
|
||||
}
|
||||
_, err = s.Store.UpsertUserSetting(ctx, userSetting)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &v1pb.Shortcut{
|
||||
Name: constructShortcutName(user.Username, foundShortcut.GetId()),
|
||||
Title: foundShortcut.GetTitle(),
|
||||
Filter: foundShortcut.GetFilter(),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) DeleteShortcut(ctx context.Context, request *v1pb.DeleteShortcutRequest) (*emptypb.Empty, error) {
|
||||
user, shortcutID, err := s.extractUserAndShortcutIDFromName(ctx, request.Name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid shortcut name: %v", err)
|
||||
}
|
||||
userID := user.ID
|
||||
|
||||
currentUser, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get current user: %v", err)
|
||||
}
|
||||
if currentUser == nil || currentUser.ID != userID {
|
||||
return nil, status.Errorf(codes.PermissionDenied, "permission denied")
|
||||
}
|
||||
|
||||
userSetting, err := s.Store.GetUserSetting(ctx, &store.FindUserSetting{
|
||||
UserID: &userID,
|
||||
Key: storepb.UserSetting_SHORTCUTS,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if userSetting == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "shortcut not found")
|
||||
}
|
||||
|
||||
shortcutsUserSetting := userSetting.GetShortcuts()
|
||||
shortcuts := shortcutsUserSetting.GetShortcuts()
|
||||
newShortcuts := make([]*storepb.ShortcutsUserSetting_Shortcut, 0, len(shortcuts))
|
||||
found := false
|
||||
for _, shortcut := range shortcuts {
|
||||
if shortcut.GetId() != shortcutID {
|
||||
newShortcuts = append(newShortcuts, shortcut)
|
||||
} else {
|
||||
found = true
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
return nil, status.Errorf(codes.NotFound, "shortcut not found")
|
||||
}
|
||||
shortcutsUserSetting.Shortcuts = newShortcuts
|
||||
userSetting.Value = &storepb.UserSetting_Shortcuts{
|
||||
Shortcuts: shortcutsUserSetting,
|
||||
}
|
||||
_, err = s.Store.UpsertUserSetting(ctx, userSetting)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to upsert user setting: %v", err)
|
||||
}
|
||||
|
||||
return &emptypb.Empty{}, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) validateFilter(ctx context.Context, filterStr string) error {
|
||||
if filterStr == "" {
|
||||
return errors.New("filter cannot be empty")
|
||||
}
|
||||
|
||||
engine, err := filter.DefaultEngine()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var dialect filter.DialectName
|
||||
switch s.Profile.Driver {
|
||||
case "mysql":
|
||||
dialect = filter.DialectMySQL
|
||||
case "postgres":
|
||||
dialect = filter.DialectPostgres
|
||||
default:
|
||||
dialect = filter.DialectSQLite
|
||||
}
|
||||
|
||||
if _, err := engine.CompileToStatement(ctx, filterStr, filter.RenderOptions{Dialect: dialect}); err != nil {
|
||||
return errors.Wrap(err, "failed to compile filter")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
package v1
|
||||
|
||||
import "github.com/usememos/memos/store"
|
||||
|
||||
func buildMemoName(uid string) string {
|
||||
return MemoNamePrefix + uid
|
||||
}
|
||||
|
||||
// resolveSSECreatorID returns the CreatorID used for SSE delivery filtering.
|
||||
// For a comment memo, it returns the parent memo's CreatorID so that private
|
||||
// parent-memo events are scoped to the parent owner.
|
||||
func resolveSSECreatorID(memo *store.Memo, parentMemo *store.Memo) int32 {
|
||||
if memo == nil {
|
||||
return 0
|
||||
}
|
||||
if parentMemo != nil {
|
||||
return parentMemo.CreatorID
|
||||
}
|
||||
return memo.CreatorID
|
||||
}
|
||||
|
||||
// buildMemoReactionSSEEvent constructs an SSEEvent for a reaction on a memo.
|
||||
// Pass parentMemo when the memo is a comment (memo.ParentUID != nil).
|
||||
func buildMemoReactionSSEEvent(eventType SSEEventType, contentID string, memo *store.Memo, parentMemo *store.Memo) *SSEEvent {
|
||||
parent := ""
|
||||
if memo != nil && memo.ParentUID != nil {
|
||||
parent = buildMemoName(*memo.ParentUID)
|
||||
}
|
||||
visibility := store.Visibility("")
|
||||
if memo != nil {
|
||||
visibility = memo.Visibility
|
||||
}
|
||||
return &SSEEvent{
|
||||
Type: eventType,
|
||||
Name: contentID,
|
||||
Parent: parent,
|
||||
Visibility: visibility,
|
||||
CreatorID: resolveSSECreatorID(memo, parentMemo),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,123 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/labstack/echo/v5"
|
||||
|
||||
"github.com/usememos/memos/server/auth"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
const (
|
||||
// sseHeartbeatInterval is the interval between heartbeat pings to keep the connection alive.
|
||||
sseHeartbeatInterval = 30 * time.Second
|
||||
)
|
||||
|
||||
type sseRouteRegistrar interface {
|
||||
GET(path string, h echo.HandlerFunc, m ...echo.MiddlewareFunc) echo.RouteInfo
|
||||
}
|
||||
|
||||
// RegisterSSERoutes registers the SSE endpoint on the given Echo router.
|
||||
func RegisterSSERoutes(router sseRouteRegistrar, hub *SSEHub, storeInstance *store.Store, secret string) {
|
||||
authenticator := auth.NewAuthenticator(storeInstance, secret)
|
||||
router.GET("/api/v1/sse", func(c *echo.Context) error {
|
||||
return handleSSE(c, hub, authenticator)
|
||||
})
|
||||
}
|
||||
|
||||
// handleSSE handles the SSE connection for live memo refresh.
|
||||
// Authentication is done via Bearer token in the Authorization header.
|
||||
func handleSSE(c *echo.Context, hub *SSEHub, authenticator *auth.Authenticator) error {
|
||||
// Authenticate the request.
|
||||
authHeader := c.Request().Header.Get("Authorization")
|
||||
result := authenticator.Authenticate(c.Request().Context(), authHeader)
|
||||
if result == nil {
|
||||
return c.JSON(http.StatusUnauthorized, map[string]string{"error": "authentication required"})
|
||||
}
|
||||
userID, role := getSSEClientIdentity(result)
|
||||
if userID == 0 {
|
||||
return c.JSON(http.StatusUnauthorized, map[string]string{"error": "authentication required"})
|
||||
}
|
||||
|
||||
// Set SSE headers.
|
||||
w := c.Response()
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
w.Header().Set("Cache-Control", "no-cache")
|
||||
w.Header().Set("Connection", "keep-alive")
|
||||
w.Header().Set("X-Accel-Buffering", "no") // Disable nginx buffering
|
||||
w.WriteHeader(http.StatusOK)
|
||||
|
||||
// Flush headers immediately.
|
||||
if f, ok := w.(http.Flusher); ok {
|
||||
f.Flush()
|
||||
}
|
||||
|
||||
// Subscribe to the hub.
|
||||
client := hub.Subscribe(userID, role)
|
||||
defer hub.Unsubscribe(client)
|
||||
|
||||
// Create a ticker for heartbeat pings.
|
||||
heartbeat := time.NewTicker(sseHeartbeatInterval)
|
||||
defer heartbeat.Stop()
|
||||
|
||||
ctx := c.Request().Context()
|
||||
|
||||
slog.Debug("SSE client connected", "userID", userID)
|
||||
|
||||
// Send an initial comment so clients and dev proxies observe the stream
|
||||
// immediately instead of waiting for the first heartbeat or data event.
|
||||
if _, err := fmt.Fprint(w, ": connected\n\n"); err != nil {
|
||||
return nil
|
||||
}
|
||||
if f, ok := w.(http.Flusher); ok {
|
||||
f.Flush()
|
||||
}
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
// Client disconnected.
|
||||
slog.Debug("SSE client disconnected", "userID", userID)
|
||||
return nil
|
||||
|
||||
case data, ok := <-client.events:
|
||||
if !ok {
|
||||
// Channel closed, client was unsubscribed.
|
||||
return nil
|
||||
}
|
||||
// Write SSE event.
|
||||
if _, err := fmt.Fprintf(w, "data: %s\n\n", data); err != nil {
|
||||
return nil
|
||||
}
|
||||
if f, ok := w.(http.Flusher); ok {
|
||||
f.Flush()
|
||||
}
|
||||
|
||||
case <-heartbeat.C:
|
||||
// Send a heartbeat comment to keep the connection alive.
|
||||
if _, err := fmt.Fprint(w, ": heartbeat\n\n"); err != nil {
|
||||
return nil
|
||||
}
|
||||
if f, ok := w.(http.Flusher); ok {
|
||||
f.Flush()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func getSSEClientIdentity(result *auth.AuthResult) (int32, store.Role) {
|
||||
if result == nil {
|
||||
return 0, store.RoleUser
|
||||
}
|
||||
if result.Claims != nil {
|
||||
return result.Claims.UserID, store.Role(result.Claims.Role)
|
||||
}
|
||||
if result.User != nil {
|
||||
return result.User.ID, result.User.Role
|
||||
}
|
||||
return 0, store.RoleUser
|
||||
}
|
||||
@@ -0,0 +1,144 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"sync"
|
||||
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
// SSEEventType represents the type of change event.
|
||||
type SSEEventType string
|
||||
|
||||
const (
|
||||
SSEEventMemoCreated SSEEventType = "memo.created"
|
||||
SSEEventMemoUpdated SSEEventType = "memo.updated"
|
||||
SSEEventMemoDeleted SSEEventType = "memo.deleted"
|
||||
SSEEventMemoCommentCreated SSEEventType = "memo.comment.created"
|
||||
SSEEventReactionUpserted SSEEventType = "reaction.upserted"
|
||||
SSEEventReactionDeleted SSEEventType = "reaction.deleted"
|
||||
)
|
||||
|
||||
// SSEEvent represents a change event sent to SSE clients.
|
||||
type SSEEvent struct {
|
||||
Type SSEEventType `json:"type"`
|
||||
// Name is the affected resource name (e.g., "memos/xxxx").
|
||||
// For reaction events, this is the memo resource name that the reaction belongs to.
|
||||
Name string `json:"name"`
|
||||
// Parent is the parent memo resource name when the affected resource is a comment.
|
||||
Parent string `json:"parent,omitempty"`
|
||||
// Visibility and CreatorID are used only for server-side delivery filtering.
|
||||
Visibility store.Visibility `json:"-"`
|
||||
CreatorID int32 `json:"-"`
|
||||
}
|
||||
|
||||
// JSON returns the JSON representation of the event.
|
||||
// Returns nil if marshaling fails (error is logged).
|
||||
func (e *SSEEvent) JSON() []byte {
|
||||
data, err := json.Marshal(e)
|
||||
if err != nil {
|
||||
slog.Error("failed to marshal SSE event", "err", err, "event", e)
|
||||
return nil
|
||||
}
|
||||
return data
|
||||
}
|
||||
|
||||
// SSEClient represents a single SSE connection.
|
||||
type SSEClient struct {
|
||||
events chan []byte
|
||||
userID int32
|
||||
role store.Role
|
||||
}
|
||||
|
||||
// SSEHub manages SSE client connections and broadcasts events.
|
||||
// It is safe for concurrent use.
|
||||
type SSEHub struct {
|
||||
mu sync.RWMutex
|
||||
clients map[*SSEClient]struct{}
|
||||
closed bool
|
||||
}
|
||||
|
||||
// NewSSEHub creates a new SSE hub.
|
||||
func NewSSEHub() *SSEHub {
|
||||
return &SSEHub{
|
||||
clients: make(map[*SSEClient]struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
// Subscribe registers a new client and returns it.
|
||||
// The caller must call Unsubscribe when done.
|
||||
func (h *SSEHub) Subscribe(userID int32, role store.Role) *SSEClient {
|
||||
c := &SSEClient{
|
||||
// Buffer a few events so a slow client doesn't block broadcasting.
|
||||
events: make(chan []byte, 32),
|
||||
userID: userID,
|
||||
role: role,
|
||||
}
|
||||
h.mu.Lock()
|
||||
if h.closed {
|
||||
close(c.events)
|
||||
} else {
|
||||
h.clients[c] = struct{}{}
|
||||
}
|
||||
h.mu.Unlock()
|
||||
return c
|
||||
}
|
||||
|
||||
// Unsubscribe removes a client and closes its channel.
|
||||
func (h *SSEHub) Unsubscribe(c *SSEClient) {
|
||||
h.mu.Lock()
|
||||
if _, ok := h.clients[c]; ok {
|
||||
delete(h.clients, c)
|
||||
close(c.events)
|
||||
}
|
||||
h.mu.Unlock()
|
||||
}
|
||||
|
||||
// Close disconnects all subscribed SSE clients.
|
||||
func (h *SSEHub) Close() {
|
||||
h.mu.Lock()
|
||||
defer h.mu.Unlock()
|
||||
if h.closed {
|
||||
return
|
||||
}
|
||||
h.closed = true
|
||||
for c := range h.clients {
|
||||
delete(h.clients, c)
|
||||
close(c.events)
|
||||
}
|
||||
}
|
||||
|
||||
// Broadcast sends an event to all connected clients.
|
||||
// Slow clients that have a full buffer will have the event dropped
|
||||
// to avoid blocking the broadcaster.
|
||||
func (h *SSEHub) Broadcast(event *SSEEvent) {
|
||||
data := event.JSON()
|
||||
if len(data) == 0 {
|
||||
return
|
||||
}
|
||||
h.mu.RLock()
|
||||
defer h.mu.RUnlock()
|
||||
for c := range h.clients {
|
||||
if !c.canReceive(event) {
|
||||
continue
|
||||
}
|
||||
select {
|
||||
case c.events <- data:
|
||||
default:
|
||||
// Drop event for slow client to avoid blocking.
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (c *SSEClient) canReceive(event *SSEEvent) bool {
|
||||
switch event.Visibility {
|
||||
case store.Private:
|
||||
return c.userID == event.CreatorID || c.role == store.RoleAdmin
|
||||
case store.Public, store.Protected, "":
|
||||
return true
|
||||
default:
|
||||
slog.Warn("SSE canReceive: unknown visibility type, denying event", "visibility", event.Visibility)
|
||||
return false
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,251 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
// helpers shared by multiple tests in this file.
|
||||
|
||||
func mustReceive(t *testing.T, ch <-chan []byte, within time.Duration) []byte {
|
||||
t.Helper()
|
||||
select {
|
||||
case data := <-ch:
|
||||
return data
|
||||
case <-time.After(within):
|
||||
t.Fatal("timed out waiting for SSE event")
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func mustNotReceive(t *testing.T, ch <-chan []byte, within time.Duration) {
|
||||
t.Helper()
|
||||
select {
|
||||
case data := <-ch:
|
||||
t.Fatalf("unexpected SSE event received: %s", data)
|
||||
case <-time.After(within):
|
||||
}
|
||||
}
|
||||
|
||||
func TestSSEHub_SubscribeUnsubscribe(t *testing.T) {
|
||||
hub := NewSSEHub()
|
||||
|
||||
client := hub.Subscribe(1, store.RoleUser)
|
||||
require.NotNil(t, client)
|
||||
require.NotNil(t, client.events)
|
||||
|
||||
// Unsubscribe removes the client and closes the channel.
|
||||
hub.Unsubscribe(client)
|
||||
|
||||
// Channel should be closed.
|
||||
_, ok := <-client.events
|
||||
assert.False(t, ok, "channel should be closed after Unsubscribe")
|
||||
}
|
||||
|
||||
func TestSSEHub_Close(t *testing.T) {
|
||||
hub := NewSSEHub()
|
||||
c1 := hub.Subscribe(1, store.RoleUser)
|
||||
c2 := hub.Subscribe(2, store.RoleAdmin)
|
||||
|
||||
hub.Close()
|
||||
hub.Close()
|
||||
|
||||
for _, ch := range []chan []byte{c1.events, c2.events} {
|
||||
_, ok := <-ch
|
||||
assert.False(t, ok, "channel should be closed after hub close")
|
||||
}
|
||||
|
||||
late := hub.Subscribe(3, store.RoleUser)
|
||||
_, ok := <-late.events
|
||||
assert.False(t, ok, "late subscriber should be closed immediately")
|
||||
|
||||
hub.Broadcast(&SSEEvent{Type: SSEEventMemoCreated, Name: "memos/123"})
|
||||
hub.Unsubscribe(c1)
|
||||
hub.Unsubscribe(late)
|
||||
}
|
||||
|
||||
func TestSSEHub_Broadcast(t *testing.T) {
|
||||
hub := NewSSEHub()
|
||||
client := hub.Subscribe(1, store.RoleUser)
|
||||
defer hub.Unsubscribe(client)
|
||||
|
||||
event := &SSEEvent{Type: SSEEventMemoCreated, Name: "memos/123"}
|
||||
hub.Broadcast(event)
|
||||
|
||||
select {
|
||||
case data := <-client.events:
|
||||
assert.Contains(t, string(data), `"type":"memo.created"`)
|
||||
assert.Contains(t, string(data), `"name":"memos/123"`)
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("expected to receive event within 1s")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSSEHub_BroadcastMultipleClients(t *testing.T) {
|
||||
hub := NewSSEHub()
|
||||
c1 := hub.Subscribe(1, store.RoleUser)
|
||||
defer hub.Unsubscribe(c1)
|
||||
c2 := hub.Subscribe(2, store.RoleUser)
|
||||
defer hub.Unsubscribe(c2)
|
||||
|
||||
event := &SSEEvent{Type: SSEEventMemoDeleted, Name: "memos/456"}
|
||||
hub.Broadcast(event)
|
||||
|
||||
for _, ch := range []chan []byte{c1.events, c2.events} {
|
||||
select {
|
||||
case data := <-ch:
|
||||
assert.Contains(t, string(data), "memo.deleted")
|
||||
assert.Contains(t, string(data), "memos/456")
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("expected to receive event within 1s")
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSSEEvent_JSON(t *testing.T) {
|
||||
e := &SSEEvent{Type: SSEEventMemoUpdated, Name: "memos/789", Parent: "memos/123"}
|
||||
data := e.JSON()
|
||||
require.NotEmpty(t, data)
|
||||
assert.Contains(t, string(data), `"type":"memo.updated"`)
|
||||
assert.Contains(t, string(data), `"name":"memos/789"`)
|
||||
assert.Contains(t, string(data), `"parent":"memos/123"`)
|
||||
}
|
||||
|
||||
func TestSSEHub_PrivateEventsAreScoped(t *testing.T) {
|
||||
hub := NewSSEHub()
|
||||
owner := hub.Subscribe(1, store.RoleUser)
|
||||
defer hub.Unsubscribe(owner)
|
||||
other := hub.Subscribe(2, store.RoleUser)
|
||||
defer hub.Unsubscribe(other)
|
||||
admin := hub.Subscribe(3, store.RoleAdmin)
|
||||
defer hub.Unsubscribe(admin)
|
||||
|
||||
hub.Broadcast(&SSEEvent{
|
||||
Type: SSEEventMemoUpdated,
|
||||
Name: "memos/private",
|
||||
Visibility: store.Private,
|
||||
CreatorID: 1,
|
||||
})
|
||||
|
||||
select {
|
||||
case <-owner.events:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("owner should receive private event")
|
||||
}
|
||||
|
||||
select {
|
||||
case <-admin.events:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("admin should receive private event")
|
||||
}
|
||||
|
||||
select {
|
||||
case <-other.events:
|
||||
t.Fatal("non-owner should not receive private event")
|
||||
case <-time.After(100 * time.Millisecond):
|
||||
}
|
||||
}
|
||||
|
||||
func TestSSEClient_CanReceive_UnknownVisibility(t *testing.T) {
|
||||
hub := NewSSEHub()
|
||||
client := hub.Subscribe(1, store.RoleUser)
|
||||
defer hub.Unsubscribe(client)
|
||||
|
||||
// An event with an unrecognised visibility value should be denied (safe default).
|
||||
hub.Broadcast(&SSEEvent{
|
||||
Type: SSEEventMemoUpdated,
|
||||
Name: "memos/unknown-vis",
|
||||
Visibility: store.Visibility("CUSTOM"),
|
||||
})
|
||||
|
||||
mustNotReceive(t, client.events, 100*time.Millisecond)
|
||||
}
|
||||
|
||||
func TestSSEHub_SlowClientEventsDropped(t *testing.T) {
|
||||
hub := NewSSEHub()
|
||||
// Subscribe but never read, so the channel fills up.
|
||||
slow := hub.Subscribe(1, store.RoleUser)
|
||||
defer hub.Unsubscribe(slow)
|
||||
|
||||
event := &SSEEvent{Type: SSEEventMemoCreated, Name: "memos/x"}
|
||||
// Send more events than the buffer capacity (32).
|
||||
for range 40 {
|
||||
hub.Broadcast(event) // must not block
|
||||
}
|
||||
|
||||
// At most 32 events should have been queued; the rest were silently dropped.
|
||||
assert.LessOrEqual(t, len(slow.events), 32)
|
||||
}
|
||||
|
||||
func TestResolveSSECreatorID(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
memo *store.Memo
|
||||
parentMemo *store.Memo
|
||||
want int32
|
||||
}{
|
||||
{
|
||||
name: "nil memo returns 0",
|
||||
memo: nil, parentMemo: nil,
|
||||
want: 0,
|
||||
},
|
||||
{
|
||||
name: "memo without parent returns memo CreatorID",
|
||||
memo: &store.Memo{CreatorID: 5},
|
||||
parentMemo: nil,
|
||||
want: 5,
|
||||
},
|
||||
{
|
||||
name: "memo with parent returns parent CreatorID",
|
||||
memo: &store.Memo{CreatorID: 5},
|
||||
parentMemo: &store.Memo{CreatorID: 9},
|
||||
want: 9,
|
||||
},
|
||||
}
|
||||
for _, tc := range tests {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
assert.Equal(t, tc.want, resolveSSECreatorID(tc.memo, tc.parentMemo))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildMemoReactionSSEEvent(t *testing.T) {
|
||||
parentUID := "parent-uid"
|
||||
|
||||
t.Run("top-level memo reaction", func(t *testing.T) {
|
||||
memo := &store.Memo{CreatorID: 10, Visibility: store.Public}
|
||||
event := buildMemoReactionSSEEvent(SSEEventReactionUpserted, "memos/abc", memo, nil)
|
||||
assert.Equal(t, SSEEventReactionUpserted, event.Type)
|
||||
assert.Equal(t, "memos/abc", event.Name)
|
||||
assert.Equal(t, "", event.Parent)
|
||||
assert.Equal(t, store.Public, event.Visibility)
|
||||
assert.Equal(t, int32(10), event.CreatorID)
|
||||
})
|
||||
|
||||
t.Run("reaction on comment is scoped to parent owner", func(t *testing.T) {
|
||||
memo := &store.Memo{
|
||||
CreatorID: 10,
|
||||
Visibility: store.Private,
|
||||
ParentUID: &parentUID,
|
||||
}
|
||||
parentMemo := &store.Memo{CreatorID: 7}
|
||||
event := buildMemoReactionSSEEvent(SSEEventReactionDeleted, "memos/abc", memo, parentMemo)
|
||||
assert.Equal(t, SSEEventReactionDeleted, event.Type)
|
||||
assert.Equal(t, MemoNamePrefix+parentUID, event.Parent)
|
||||
assert.Equal(t, store.Private, event.Visibility)
|
||||
assert.Equal(t, int32(7), event.CreatorID)
|
||||
})
|
||||
|
||||
t.Run("nil memo produces a safe zero-value event", func(t *testing.T) {
|
||||
event := buildMemoReactionSSEEvent(SSEEventReactionUpserted, "memos/abc", nil, nil)
|
||||
assert.Equal(t, "memos/abc", event.Name)
|
||||
assert.Equal(t, "", event.Parent)
|
||||
assert.Equal(t, store.Visibility(""), event.Visibility)
|
||||
assert.Equal(t, int32(0), event.CreatorID)
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,314 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/usememos/memos/internal/profile"
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/server/auth"
|
||||
"github.com/usememos/memos/store"
|
||||
teststore "github.com/usememos/memos/store/test"
|
||||
)
|
||||
|
||||
// newIntegrationService builds a minimal APIV1Service backed by an in-memory
|
||||
// SQLite database. The store is closed automatically via t.Cleanup.
|
||||
func newIntegrationService(t *testing.T) *APIV1Service {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
st := teststore.NewTestingStore(ctx, t)
|
||||
t.Cleanup(func() { st.Close() })
|
||||
p := &profile.Profile{Demo: true, Data: t.TempDir(), Driver: "sqlite", DSN: ":memory:"}
|
||||
return NewAPIV1Service("test-secret", p, st)
|
||||
}
|
||||
|
||||
// userCtx returns a context that authenticates as the given user.
|
||||
func userCtx(ctx context.Context, userID int32) context.Context {
|
||||
return context.WithValue(ctx, auth.UserIDContextKey, userID)
|
||||
}
|
||||
|
||||
// collectEventsFor reads events from ch for the given duration and returns them.
|
||||
func collectEventsFor(ch <-chan []byte, d time.Duration) []string {
|
||||
var out []string
|
||||
deadline := time.After(d)
|
||||
for {
|
||||
select {
|
||||
case data := <-ch:
|
||||
out = append(out, string(data))
|
||||
case <-deadline:
|
||||
return out
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---- context suppression ----
|
||||
|
||||
func TestSuppressSSEContext(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("default context is not suppressed", func(t *testing.T) {
|
||||
assert.False(t, isSSESuppressed(ctx))
|
||||
})
|
||||
|
||||
t.Run("withSuppressSSE marks context as suppressed", func(t *testing.T) {
|
||||
assert.True(t, isSSESuppressed(withSuppressSSE(ctx)))
|
||||
})
|
||||
|
||||
t.Run("suppression does not bleed into parent context", func(t *testing.T) {
|
||||
suppressed := withSuppressSSE(ctx)
|
||||
_ = suppressed
|
||||
assert.False(t, isSSESuppressed(ctx))
|
||||
})
|
||||
}
|
||||
|
||||
// ---- CreateMemoComment double-broadcast fix ----
|
||||
|
||||
func TestCreateMemoComment_NoDuplicateSSEBroadcast(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc := newIntegrationService(t)
|
||||
|
||||
// Create an admin so the store is initialised, then a regular commenter.
|
||||
author, err := svc.Store.CreateUser(ctx, &store.User{
|
||||
Username: "author", Role: store.RoleAdmin, Email: "author@example.com",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
commenter, err := svc.Store.CreateUser(ctx, &store.User{
|
||||
Username: "commenter", Role: store.RoleUser, Email: "commenter@example.com",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
authorCtx := userCtx(ctx, author.ID)
|
||||
commenterCtx := userCtx(ctx, commenter.ID)
|
||||
|
||||
// Create a public memo so the commenter can react.
|
||||
parent, err := svc.CreateMemo(authorCtx, &v1pb.CreateMemoRequest{
|
||||
Memo: &v1pb.Memo{Content: "parent memo", Visibility: v1pb.Visibility_PUBLIC},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Subscribe after the parent memo is created so the memo.created event
|
||||
// for the parent does not pollute the assertion window.
|
||||
client := svc.SSEHub.Subscribe(author.ID, store.RoleAdmin)
|
||||
defer svc.SSEHub.Unsubscribe(client)
|
||||
|
||||
// Create a comment. Before the fix, this fired both memo.created (for the
|
||||
// comment memo) and memo.comment.created (for the parent).
|
||||
_, err = svc.CreateMemoComment(commenterCtx, &v1pb.CreateMemoCommentRequest{
|
||||
Name: parent.Name,
|
||||
Comment: &v1pb.Memo{Content: "a comment", Visibility: v1pb.Visibility_PUBLIC},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Give the synchronous broadcast a moment to land in the buffer, then
|
||||
// collect everything that arrived.
|
||||
events := collectEventsFor(client.events, 150*time.Millisecond)
|
||||
|
||||
require.Len(t, events, 1, "expected exactly one SSE event for a comment creation, got: %v", events)
|
||||
assert.True(t, strings.Contains(events[0], `"memo.comment.created"`),
|
||||
"expected memo.comment.created, got: %s", events[0])
|
||||
}
|
||||
|
||||
func TestCreateMemoWithAttachment_NoDuplicateUpdatedSSEBroadcast(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc := newIntegrationService(t)
|
||||
|
||||
user, err := svc.Store.CreateUser(ctx, &store.User{
|
||||
Username: "user", Role: store.RoleAdmin, Email: "user@example.com",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
uctx := userCtx(ctx, user.ID)
|
||||
|
||||
attachment, err := svc.CreateAttachment(uctx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{
|
||||
Filename: "test.txt",
|
||||
Size: 5,
|
||||
Type: "text/plain",
|
||||
Content: []byte("hello"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
client := svc.SSEHub.Subscribe(user.ID, store.RoleAdmin)
|
||||
defer svc.SSEHub.Unsubscribe(client)
|
||||
|
||||
memo, err := svc.CreateMemo(uctx, &v1pb.CreateMemoRequest{
|
||||
Memo: &v1pb.Memo{
|
||||
Content: "memo with initial attachment",
|
||||
Visibility: v1pb.Visibility_PUBLIC,
|
||||
Attachments: []*v1pb.Attachment{
|
||||
{Name: attachment.Name},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
events := collectEventsFor(client.events, 150*time.Millisecond)
|
||||
|
||||
require.Len(t, events, 1, "expected exactly one SSE event for memo creation with attachment, got: %v", events)
|
||||
assert.Contains(t, events[0], `"memo.created"`)
|
||||
assert.Contains(t, events[0], memo.Name)
|
||||
assert.NotContains(t, events[0], `"memo.updated"`)
|
||||
}
|
||||
|
||||
// ---- Reaction SSE events carry correct visibility / parent fields ----
|
||||
|
||||
func TestUpsertMemoReaction_SSEEvent(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc := newIntegrationService(t)
|
||||
|
||||
user, err := svc.Store.CreateUser(ctx, &store.User{
|
||||
Username: "user", Role: store.RoleAdmin, Email: "user@example.com",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
uctx := userCtx(ctx, user.ID)
|
||||
|
||||
memo, err := svc.CreateMemo(uctx, &v1pb.CreateMemoRequest{
|
||||
Memo: &v1pb.Memo{Content: "reacted memo", Visibility: v1pb.Visibility_PUBLIC},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
client := svc.SSEHub.Subscribe(user.ID, store.RoleAdmin)
|
||||
defer svc.SSEHub.Unsubscribe(client)
|
||||
|
||||
_, err = svc.UpsertMemoReaction(uctx, &v1pb.UpsertMemoReactionRequest{
|
||||
Name: memo.Name,
|
||||
Reaction: &v1pb.Reaction{
|
||||
ContentId: memo.Name,
|
||||
ReactionType: "👍",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
data := mustReceive(t, client.events, time.Second)
|
||||
payload := string(data)
|
||||
assert.Contains(t, payload, `"reaction.upserted"`)
|
||||
assert.Contains(t, payload, memo.Name)
|
||||
mustNotReceive(t, client.events, 100*time.Millisecond)
|
||||
}
|
||||
|
||||
func TestDeleteMemoReaction_SSEEvent(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc := newIntegrationService(t)
|
||||
|
||||
user, err := svc.Store.CreateUser(ctx, &store.User{
|
||||
Username: "user", Role: store.RoleAdmin, Email: "user@example.com",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
uctx := userCtx(ctx, user.ID)
|
||||
|
||||
memo, err := svc.CreateMemo(uctx, &v1pb.CreateMemoRequest{
|
||||
Memo: &v1pb.Memo{Content: "reacted memo", Visibility: v1pb.Visibility_PUBLIC},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
reaction, err := svc.UpsertMemoReaction(uctx, &v1pb.UpsertMemoReactionRequest{
|
||||
Name: memo.Name,
|
||||
Reaction: &v1pb.Reaction{
|
||||
ContentId: memo.Name,
|
||||
ReactionType: "❤️",
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
client := svc.SSEHub.Subscribe(user.ID, store.RoleAdmin)
|
||||
defer svc.SSEHub.Unsubscribe(client)
|
||||
|
||||
_, err = svc.DeleteMemoReaction(uctx, &v1pb.DeleteMemoReactionRequest{
|
||||
Name: reaction.Name,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
data := mustReceive(t, client.events, time.Second)
|
||||
payload := string(data)
|
||||
assert.Contains(t, payload, `"reaction.deleted"`)
|
||||
assert.Contains(t, payload, memo.Name)
|
||||
mustNotReceive(t, client.events, 100*time.Millisecond)
|
||||
}
|
||||
|
||||
func TestSetMemoAttachments_EmitsMemoUpdatedSSEEvent(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc := newIntegrationService(t)
|
||||
|
||||
user, err := svc.Store.CreateUser(ctx, &store.User{
|
||||
Username: "user", Role: store.RoleAdmin, Email: "user@example.com",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
uctx := userCtx(ctx, user.ID)
|
||||
|
||||
memo, err := svc.CreateMemo(uctx, &v1pb.CreateMemoRequest{
|
||||
Memo: &v1pb.Memo{Content: "memo with attachments", Visibility: v1pb.Visibility_PUBLIC},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
attachment, err := svc.CreateAttachment(uctx, &v1pb.CreateAttachmentRequest{
|
||||
Attachment: &v1pb.Attachment{
|
||||
Filename: "test.txt",
|
||||
Size: 5,
|
||||
Type: "text/plain",
|
||||
Content: []byte("hello"),
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
client := svc.SSEHub.Subscribe(user.ID, store.RoleAdmin)
|
||||
defer svc.SSEHub.Unsubscribe(client)
|
||||
|
||||
_, err = svc.SetMemoAttachments(uctx, &v1pb.SetMemoAttachmentsRequest{
|
||||
Name: memo.Name,
|
||||
Attachments: []*v1pb.Attachment{
|
||||
{Name: attachment.Name},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
data := mustReceive(t, client.events, time.Second)
|
||||
payload := string(data)
|
||||
assert.Contains(t, payload, `"memo.updated"`)
|
||||
assert.Contains(t, payload, memo.Name)
|
||||
mustNotReceive(t, client.events, 100*time.Millisecond)
|
||||
}
|
||||
|
||||
func TestSetMemoRelations_EmitsMemoUpdatedSSEEvent(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
svc := newIntegrationService(t)
|
||||
|
||||
user, err := svc.Store.CreateUser(ctx, &store.User{
|
||||
Username: "user", Role: store.RoleAdmin, Email: "user@example.com",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
uctx := userCtx(ctx, user.ID)
|
||||
|
||||
memo1, err := svc.CreateMemo(uctx, &v1pb.CreateMemoRequest{
|
||||
Memo: &v1pb.Memo{Content: "memo one", Visibility: v1pb.Visibility_PUBLIC},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
memo2, err := svc.CreateMemo(uctx, &v1pb.CreateMemoRequest{
|
||||
Memo: &v1pb.Memo{Content: "memo two", Visibility: v1pb.Visibility_PUBLIC},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
client := svc.SSEHub.Subscribe(user.ID, store.RoleAdmin)
|
||||
defer svc.SSEHub.Unsubscribe(client)
|
||||
|
||||
_, err = svc.SetMemoRelations(uctx, &v1pb.SetMemoRelationsRequest{
|
||||
Name: memo1.Name,
|
||||
Relations: []*v1pb.MemoRelation{
|
||||
{
|
||||
RelatedMemo: &v1pb.MemoRelation_Memo{Name: memo2.Name},
|
||||
Type: v1pb.MemoRelation_REFERENCE,
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
data := mustReceive(t, client.events, time.Second)
|
||||
payload := string(data)
|
||||
assert.Contains(t, payload, `"memo.updated"`)
|
||||
assert.Contains(t, payload, memo1.Name)
|
||||
mustNotReceive(t, client.events, 100*time.Millisecond)
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/usememos/memos/internal/util"
|
||||
)
|
||||
|
||||
// deriveSSOUsername produces the local username for a new SSO-created user.
|
||||
//
|
||||
// The current policy is to use a standard UUID string directly. This keeps the
|
||||
// username independent of IdP profile fields and avoids availability probes or
|
||||
// retry loops around concurrent first-time logins.
|
||||
func deriveSSOUsername() (string, error) {
|
||||
username := util.GenUUID()
|
||||
if err := validateWritableUsername(username); err != nil {
|
||||
return "", errors.Wrap(err, "generated UUID did not satisfy username constraints")
|
||||
}
|
||||
return username, nil
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/usememos/memos/internal/base"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
// BuildUserName returns the canonical public resource name for a user.
|
||||
func BuildUserName(username string) string {
|
||||
return UserNamePrefix + username
|
||||
}
|
||||
|
||||
func parseUsernameFromName(name string) (string, error) {
|
||||
tokens, err := GetNameParentTokens(name, UserNamePrefix)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
username := tokens[0]
|
||||
if username == "" {
|
||||
return "", errors.Errorf("invalid user name %q", name)
|
||||
}
|
||||
return username, nil
|
||||
}
|
||||
|
||||
func validateWritableUsername(username string) error {
|
||||
if username == "" || isNumericUsername(username) || !base.UIDMatcher.MatchString(username) {
|
||||
return errors.Errorf("invalid username %q", username)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func isNumericUsername(username string) bool {
|
||||
if username == "" {
|
||||
return false
|
||||
}
|
||||
for _, char := range username {
|
||||
if char < '0' || char > '9' {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// ResolveUserByName resolves a username-based user resource name to a store user.
|
||||
func ResolveUserByName(ctx context.Context, stores *store.Store, name string) (*store.User, error) {
|
||||
username, err := parseUsernameFromName(name)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
user, err := stores.GetUser(ctx, &store.FindUser{Username: &username})
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "resolve user by name: GetUser failed")
|
||||
}
|
||||
return user, nil
|
||||
}
|
||||
@@ -0,0 +1,116 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestValidateWritableUsername(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
username string
|
||||
wantError bool
|
||||
}{
|
||||
{
|
||||
name: "lowercase",
|
||||
username: "alice",
|
||||
},
|
||||
{
|
||||
name: "mixed case",
|
||||
username: "Alice",
|
||||
},
|
||||
{
|
||||
name: "hyphenated",
|
||||
username: "alice-smith",
|
||||
},
|
||||
{
|
||||
name: "uuid",
|
||||
username: "550e8400-e29b-41d4-a716-446655440000",
|
||||
},
|
||||
{
|
||||
name: "empty",
|
||||
username: "",
|
||||
wantError: true,
|
||||
},
|
||||
{
|
||||
name: "numeric",
|
||||
username: "123",
|
||||
wantError: true,
|
||||
},
|
||||
{
|
||||
name: "email",
|
||||
username: "alice@example.com",
|
||||
wantError: true,
|
||||
},
|
||||
{
|
||||
name: "underscore",
|
||||
username: "alice_smith",
|
||||
wantError: true,
|
||||
},
|
||||
{
|
||||
name: "space",
|
||||
username: "alice smith",
|
||||
wantError: true,
|
||||
},
|
||||
{
|
||||
name: "slash",
|
||||
username: "alice/smith",
|
||||
wantError: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
err := validateWritableUsername(test.username)
|
||||
if test.wantError && err == nil {
|
||||
t.Fatalf("validateWritableUsername(%q) succeeded, want error", test.username)
|
||||
}
|
||||
if !test.wantError && err != nil {
|
||||
t.Fatalf("validateWritableUsername(%q) returned error: %v", test.username, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseUsernameFromNameAllowsLegacyUsernames(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
want string
|
||||
wantFail bool
|
||||
}{
|
||||
{
|
||||
name: "users/alice",
|
||||
want: "alice",
|
||||
},
|
||||
{
|
||||
name: "users/alice@example.com",
|
||||
want: "alice@example.com",
|
||||
},
|
||||
{
|
||||
name: "users/alice_smith",
|
||||
want: "alice_smith",
|
||||
},
|
||||
{
|
||||
name: "users/",
|
||||
wantFail: true,
|
||||
},
|
||||
{
|
||||
name: "invalid/alice",
|
||||
wantFail: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
got, err := parseUsernameFromName(test.name)
|
||||
if test.wantFail && err == nil {
|
||||
t.Fatalf("parseUsernameFromName(%q) succeeded, want error", test.name)
|
||||
}
|
||||
if !test.wantFail && err != nil {
|
||||
t.Fatalf("parseUsernameFromName(%q) returned error: %v", test.name, err)
|
||||
}
|
||||
if got != test.want {
|
||||
t.Fatalf("parseUsernameFromName(%q) = %q, want %q", test.name, got, test.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,296 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
"google.golang.org/protobuf/types/known/timestamppb"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func (s *APIV1Service) listUsersByID(ctx context.Context, userIDs []int32) (map[int32]*store.User, error) {
|
||||
if len(userIDs) == 0 {
|
||||
return map[int32]*store.User{}, nil
|
||||
}
|
||||
|
||||
uniqueUserIDs := make([]int32, 0, len(userIDs))
|
||||
seenUserIDs := make(map[int32]struct{}, len(userIDs))
|
||||
for _, userID := range userIDs {
|
||||
if _, seen := seenUserIDs[userID]; seen {
|
||||
continue
|
||||
}
|
||||
seenUserIDs[userID] = struct{}{}
|
||||
uniqueUserIDs = append(uniqueUserIDs, userID)
|
||||
}
|
||||
|
||||
users, err := s.Store.ListUsers(ctx, &store.FindUser{IDList: uniqueUserIDs})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
usersByID := make(map[int32]*store.User, len(users))
|
||||
for _, user := range users {
|
||||
usersByID[user.ID] = user
|
||||
}
|
||||
return usersByID, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) listUsernamesByID(ctx context.Context, userIDs []int32) (map[int32]string, error) {
|
||||
usersByID, err := s.listUsersByID(ctx, userIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
usernamesByID := make(map[int32]string, len(usersByID))
|
||||
for _, user := range usersByID {
|
||||
usernamesByID[user.ID] = user.Username
|
||||
}
|
||||
return usernamesByID, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) ListAllUserStats(ctx context.Context, request *v1pb.ListAllUserStatsRequest) (*v1pb.ListAllUserStatsResponse, error) {
|
||||
rowStatus := convertStateToStore(request.State)
|
||||
memoFind := &store.FindMemo{
|
||||
// Exclude comments by default.
|
||||
ExcludeComments: true,
|
||||
ExcludeContent: true,
|
||||
RowStatus: &rowStatus,
|
||||
}
|
||||
|
||||
currentUser, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get user: %v", err)
|
||||
}
|
||||
|
||||
if request.Filter != "" {
|
||||
if err := s.validateFilter(ctx, request.Filter); err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid filter: %v", err)
|
||||
}
|
||||
memoFind.Filters = append(memoFind.Filters, request.Filter)
|
||||
}
|
||||
|
||||
if request.State == v1pb.State_ARCHIVED {
|
||||
// Archived memos are only visible to their creator.
|
||||
if currentUser == nil {
|
||||
return &v1pb.ListAllUserStatsResponse{}, nil
|
||||
}
|
||||
memoFind.CreatorID = ¤tUser.ID
|
||||
} else if currentUser == nil {
|
||||
memoFind.VisibilityList = []store.Visibility{store.Public}
|
||||
} else {
|
||||
if memoFind.CreatorID == nil {
|
||||
filter := fmt.Sprintf(`creator_id == %d || visibility in ["PUBLIC", "PROTECTED"]`, currentUser.ID)
|
||||
memoFind.Filters = append(memoFind.Filters, filter)
|
||||
} else if *memoFind.CreatorID != currentUser.ID {
|
||||
memoFind.VisibilityList = []store.Visibility{store.Public, store.Protected}
|
||||
}
|
||||
}
|
||||
|
||||
userMemoStatMap := make(map[int32]*v1pb.UserStats)
|
||||
pinnedMemoIDsByUserID := make(map[int32][]int32)
|
||||
limit := 1000
|
||||
offset := 0
|
||||
memoFind.Limit = &limit
|
||||
memoFind.Offset = &offset
|
||||
|
||||
for {
|
||||
memos, err := s.Store.ListMemos(ctx, memoFind)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to list memos: %v", err)
|
||||
}
|
||||
if len(memos) == 0 {
|
||||
break
|
||||
}
|
||||
|
||||
for _, memo := range memos {
|
||||
// Initialize user stats if not exists
|
||||
if _, exists := userMemoStatMap[memo.CreatorID]; !exists {
|
||||
userMemoStatMap[memo.CreatorID] = &v1pb.UserStats{
|
||||
Name: "",
|
||||
TagCount: make(map[string]int32),
|
||||
MemoCreatedTimestamps: []*timestamppb.Timestamp{},
|
||||
MemoUpdatedTimestamps: []*timestamppb.Timestamp{},
|
||||
PinnedMemos: []string{},
|
||||
MemoTypeStats: &v1pb.UserStats_MemoTypeStats{
|
||||
LinkCount: 0,
|
||||
CodeCount: 0,
|
||||
TodoCount: 0,
|
||||
UndoCount: 0,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
stats := userMemoStatMap[memo.CreatorID]
|
||||
|
||||
stats.MemoCreatedTimestamps = append(stats.MemoCreatedTimestamps, timestamppb.New(time.Unix(memo.CreatedTs, 0)))
|
||||
stats.MemoUpdatedTimestamps = append(stats.MemoUpdatedTimestamps, timestamppb.New(time.Unix(memo.UpdatedTs, 0)))
|
||||
|
||||
// Count memo stats
|
||||
stats.TotalMemoCount++
|
||||
|
||||
// Count tags and other properties
|
||||
if memo.Payload != nil {
|
||||
for _, tag := range memo.Payload.Tags {
|
||||
stats.TagCount[tag]++
|
||||
}
|
||||
if memo.Payload.Property != nil {
|
||||
if memo.Payload.Property.HasLink {
|
||||
stats.MemoTypeStats.LinkCount++
|
||||
}
|
||||
if memo.Payload.Property.HasCode {
|
||||
stats.MemoTypeStats.CodeCount++
|
||||
}
|
||||
if memo.Payload.Property.HasTaskList {
|
||||
stats.MemoTypeStats.TodoCount++
|
||||
}
|
||||
if memo.Payload.Property.HasIncompleteTasks {
|
||||
stats.MemoTypeStats.UndoCount++
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Track pinned memos
|
||||
if memo.Pinned {
|
||||
pinnedMemoIDsByUserID[memo.CreatorID] = append(pinnedMemoIDsByUserID[memo.CreatorID], memo.ID)
|
||||
}
|
||||
}
|
||||
|
||||
offset += limit
|
||||
}
|
||||
|
||||
userMemoStats := []*v1pb.UserStats{}
|
||||
userIDs := make([]int32, 0, len(userMemoStatMap))
|
||||
for userID := range userMemoStatMap {
|
||||
userIDs = append(userIDs, userID)
|
||||
}
|
||||
usernamesByID, err := s.listUsernamesByID(ctx, userIDs)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to list users: %v", err)
|
||||
}
|
||||
for userID, userMemoStat := range userMemoStatMap {
|
||||
username, ok := usernamesByID[userID]
|
||||
if !ok {
|
||||
return nil, status.Errorf(codes.Internal, "failed to resolve user stats name")
|
||||
}
|
||||
userMemoStat.Name = fmt.Sprintf("%s/stats", BuildUserName(username))
|
||||
for _, memoID := range pinnedMemoIDsByUserID[userID] {
|
||||
userMemoStat.PinnedMemos = append(userMemoStat.PinnedMemos, fmt.Sprintf("%s/memos/%d", BuildUserName(username), memoID))
|
||||
}
|
||||
userMemoStats = append(userMemoStats, userMemoStat)
|
||||
}
|
||||
|
||||
response := &v1pb.ListAllUserStatsResponse{
|
||||
Stats: userMemoStats,
|
||||
}
|
||||
return response, nil
|
||||
}
|
||||
|
||||
func (s *APIV1Service) GetUserStats(ctx context.Context, request *v1pb.GetUserStatsRequest) (*v1pb.UserStats, error) {
|
||||
user, err := ResolveUserByName(ctx, s.Store, request.Name)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.InvalidArgument, "invalid user name: %v", err)
|
||||
}
|
||||
if user == nil {
|
||||
return nil, status.Errorf(codes.NotFound, "user not found")
|
||||
}
|
||||
userID := user.ID
|
||||
|
||||
currentUser, err := s.fetchCurrentUser(ctx)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to get user: %v", err)
|
||||
}
|
||||
|
||||
normalStatus := store.Normal
|
||||
memoFind := &store.FindMemo{
|
||||
CreatorID: &userID,
|
||||
// Exclude comments by default.
|
||||
ExcludeComments: true,
|
||||
ExcludeContent: true,
|
||||
RowStatus: &normalStatus,
|
||||
}
|
||||
|
||||
if currentUser == nil {
|
||||
memoFind.VisibilityList = []store.Visibility{store.Public}
|
||||
} else if currentUser.ID != userID {
|
||||
memoFind.VisibilityList = []store.Visibility{store.Public, store.Protected}
|
||||
}
|
||||
|
||||
createdTimestamps := []*timestamppb.Timestamp{}
|
||||
updatedTimestamps := []*timestamppb.Timestamp{}
|
||||
tagCount := make(map[string]int32)
|
||||
linkCount := int32(0)
|
||||
codeCount := int32(0)
|
||||
todoCount := int32(0)
|
||||
undoCount := int32(0)
|
||||
pinnedMemos := []string{}
|
||||
totalMemoCount := int32(0)
|
||||
|
||||
limit := 1000
|
||||
offset := 0
|
||||
memoFind.Limit = &limit
|
||||
memoFind.Offset = &offset
|
||||
|
||||
for {
|
||||
memos, err := s.Store.ListMemos(ctx, memoFind)
|
||||
if err != nil {
|
||||
return nil, status.Errorf(codes.Internal, "failed to list memos: %v", err)
|
||||
}
|
||||
if len(memos) == 0 {
|
||||
break
|
||||
}
|
||||
|
||||
totalMemoCount += int32(len(memos))
|
||||
|
||||
for _, memo := range memos {
|
||||
createdTimestamps = append(createdTimestamps, timestamppb.New(time.Unix(memo.CreatedTs, 0)))
|
||||
updatedTimestamps = append(updatedTimestamps, timestamppb.New(time.Unix(memo.UpdatedTs, 0)))
|
||||
// Count different memo types based on content.
|
||||
if memo.Payload != nil {
|
||||
for _, tag := range memo.Payload.Tags {
|
||||
tagCount[tag]++
|
||||
}
|
||||
if memo.Payload.Property != nil {
|
||||
if memo.Payload.Property.HasLink {
|
||||
linkCount++
|
||||
}
|
||||
if memo.Payload.Property.HasCode {
|
||||
codeCount++
|
||||
}
|
||||
if memo.Payload.Property.HasTaskList {
|
||||
todoCount++
|
||||
}
|
||||
if memo.Payload.Property.HasIncompleteTasks {
|
||||
undoCount++
|
||||
}
|
||||
}
|
||||
}
|
||||
if memo.Pinned {
|
||||
pinnedMemos = append(pinnedMemos, fmt.Sprintf("%s/memos/%d", BuildUserName(user.Username), memo.ID))
|
||||
}
|
||||
}
|
||||
|
||||
offset += limit
|
||||
}
|
||||
|
||||
userStats := &v1pb.UserStats{
|
||||
Name: fmt.Sprintf("%s/stats", BuildUserName(user.Username)),
|
||||
MemoCreatedTimestamps: createdTimestamps,
|
||||
MemoUpdatedTimestamps: updatedTimestamps,
|
||||
TagCount: tagCount,
|
||||
PinnedMemos: pinnedMemos,
|
||||
TotalMemoCount: totalMemoCount,
|
||||
MemoTypeStats: &v1pb.UserStats_MemoTypeStats{
|
||||
LinkCount: linkCount,
|
||||
CodeCount: codeCount,
|
||||
TodoCount: todoCount,
|
||||
UndoCount: undoCount,
|
||||
},
|
||||
}
|
||||
|
||||
return userStats, nil
|
||||
}
|
||||
@@ -0,0 +1,150 @@
|
||||
package v1
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
|
||||
"connectrpc.com/connect"
|
||||
"github.com/grpc-ecosystem/grpc-gateway/v2/runtime"
|
||||
"github.com/labstack/echo/v5"
|
||||
"golang.org/x/sync/semaphore"
|
||||
|
||||
"github.com/usememos/memos/internal/markdown"
|
||||
"github.com/usememos/memos/internal/profile"
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/server/auth"
|
||||
"github.com/usememos/memos/server/notification"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
const maxAPIRequestBytes = 256 << 20
|
||||
|
||||
type APIV1Service struct {
|
||||
v1pb.UnimplementedInstanceServiceServer
|
||||
v1pb.UnimplementedAuthServiceServer
|
||||
v1pb.UnimplementedUserServiceServer
|
||||
v1pb.UnimplementedMemoServiceServer
|
||||
v1pb.UnimplementedAttachmentServiceServer
|
||||
v1pb.UnimplementedAIServiceServer
|
||||
v1pb.UnimplementedShortcutServiceServer
|
||||
v1pb.UnimplementedIdentityProviderServiceServer
|
||||
|
||||
Secret string
|
||||
Profile *profile.Profile
|
||||
Store *store.Store
|
||||
MarkdownService markdown.Service
|
||||
SSEHub *SSEHub
|
||||
NotificationEmailSender notification.EmailSender
|
||||
|
||||
// thumbnailSemaphore limits concurrent thumbnail generation to prevent memory exhaustion
|
||||
thumbnailSemaphore *semaphore.Weighted
|
||||
imageProcessingSemaphore *semaphore.Weighted
|
||||
|
||||
// instanceStatsCache memoizes GetInstanceStats results for instanceStatsCacheTTL.
|
||||
instanceStatsCache instanceStatsCache
|
||||
}
|
||||
|
||||
func NewAPIV1Service(secret string, profile *profile.Profile, store *store.Store) *APIV1Service {
|
||||
markdownService := markdown.NewService(
|
||||
markdown.WithTagExtension(),
|
||||
markdown.WithMentionExtension(),
|
||||
)
|
||||
return &APIV1Service{
|
||||
Secret: secret,
|
||||
Profile: profile,
|
||||
Store: store,
|
||||
MarkdownService: markdownService,
|
||||
SSEHub: NewSSEHub(),
|
||||
NotificationEmailSender: nil,
|
||||
thumbnailSemaphore: semaphore.NewWeighted(3), // Limit to 3 concurrent thumbnail generations
|
||||
imageProcessingSemaphore: semaphore.NewWeighted(2),
|
||||
}
|
||||
}
|
||||
|
||||
// RegisterGateway registers the gRPC-Gateway and Connect handlers with the given Echo instance.
|
||||
func (s *APIV1Service) RegisterGateway(ctx context.Context, echoServer *echo.Echo) error {
|
||||
// Auth middleware for gRPC-Gateway - runs after routing, has access to method name.
|
||||
// Uses the same PublicMethods config as the Connect AuthInterceptor.
|
||||
authenticator := auth.NewAuthenticator(s.Store, s.Secret)
|
||||
gatewayAuthMiddleware := func(next runtime.HandlerFunc) runtime.HandlerFunc {
|
||||
return func(w http.ResponseWriter, r *http.Request, pathParams map[string]string) {
|
||||
ctx := r.Context()
|
||||
|
||||
// Get the RPC method name from context (set by grpc-gateway after routing)
|
||||
rpcMethod, ok := runtime.RPCMethod(ctx)
|
||||
|
||||
// Extract credentials from HTTP headers
|
||||
authHeader := r.Header.Get("Authorization")
|
||||
|
||||
result := authenticator.Authenticate(ctx, authHeader)
|
||||
|
||||
// Enforce authentication for non-public methods
|
||||
// If rpcMethod cannot be determined, allow through, service layer will handle visibility checks
|
||||
if result == nil && ok && !IsPublicMethod(rpcMethod) {
|
||||
http.Error(w, `{"code": 16, "message": "authentication required"}`, http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
|
||||
// Apply auth result to context (no-op when result is nil for public endpoints)
|
||||
if result != nil {
|
||||
ctx = auth.ApplyToContext(ctx, result)
|
||||
r = r.WithContext(ctx)
|
||||
}
|
||||
|
||||
next(w, r, pathParams)
|
||||
}
|
||||
}
|
||||
|
||||
// Create gRPC-Gateway mux with auth middleware.
|
||||
gwMux := runtime.NewServeMux(
|
||||
runtime.WithMiddlewares(gatewayAuthMiddleware),
|
||||
)
|
||||
if err := v1pb.RegisterInstanceServiceHandlerServer(ctx, gwMux, s); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := v1pb.RegisterAuthServiceHandlerServer(ctx, gwMux, s); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := v1pb.RegisterUserServiceHandlerServer(ctx, gwMux, s); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := v1pb.RegisterMemoServiceHandlerServer(ctx, gwMux, s); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := v1pb.RegisterAttachmentServiceHandlerServer(ctx, gwMux, s); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := v1pb.RegisterAIServiceHandlerServer(ctx, gwMux, s); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := v1pb.RegisterShortcutServiceHandlerServer(ctx, gwMux, s); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := v1pb.RegisterIdentityProviderServiceHandlerServer(ctx, gwMux, s); err != nil {
|
||||
return err
|
||||
}
|
||||
gwGroup := echoServer.Group("")
|
||||
// Register SSE endpoint with same CORS as rest of /api/v1.
|
||||
RegisterSSERoutes(gwGroup, s.SSEHub, s.Store, s.Secret)
|
||||
handler := echo.WrapHandler(http.MaxBytesHandler(gwMux, maxAPIRequestBytes))
|
||||
|
||||
gwGroup.Any("/api/v1/*", handler)
|
||||
gwGroup.Any("/file/*", handler)
|
||||
|
||||
// Connect handlers for browser clients (replaces grpc-web).
|
||||
logStacktraces := s.Profile.Demo
|
||||
connectInterceptors := connect.WithInterceptors(
|
||||
NewMetadataInterceptor(), // Convert HTTP headers to gRPC metadata first
|
||||
NewLoggingInterceptor(logStacktraces),
|
||||
NewRecoveryInterceptor(logStacktraces),
|
||||
NewAuthInterceptor(s.Store, s.Secret),
|
||||
)
|
||||
connectMux := http.NewServeMux()
|
||||
connectHandler := NewConnectServiceHandler(s)
|
||||
connectHandler.RegisterConnectHandlers(connectMux, connectInterceptors, connect.WithReadMaxBytes(maxAPIRequestBytes))
|
||||
|
||||
connectGroup := echoServer.Group("")
|
||||
connectGroup.Any("/memos.api.v1.*", echo.WrapHandler(http.MaxBytesHandler(connectMux, maxAPIRequestBytes)))
|
||||
|
||||
return nil
|
||||
}
|
||||
Reference in New Issue
Block a user