Compare commits

...

1 Commits

Author SHA1 Message Date
johnnyjoygh 5c5d866a5b fix(auth): provision SSO users atomically 2026-07-25 01:12:15 +08:00
14 changed files with 774 additions and 99 deletions
+8 -8
View File
@@ -110,11 +110,11 @@ func (s *APIV1Service) SignIn(ctx context.Context, request *v1pb.SignInRequest)
// 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-backed 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.
// Lookup goes through the user_identity table instead of using userInfo.Identifier
// as the local lookup key. On the miss path, a local user is created with the
// identifier as its username when valid and available, or a UUID fallback
// otherwise, and the (provider, extern_uid) linkage is committed atomically with
// the user. When currentUser is provided by a caller outside AuthService.SignIn,
// the lookup miss path binds the external identity to that existing user instead.
// Concurrent first logins reconcile uniqueness conflicts by loading the linkage
// winner.
+82 -45
View File
@@ -5,6 +5,7 @@ import (
"regexp"
"strings"
"github.com/pkg/errors"
"golang.org/x/crypto/bcrypt"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
@@ -19,6 +20,11 @@ import (
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
// Defense in depth: an empty subject must never key a lookup or provision an
// account, regardless of whether the IdP layer already rejected it.
if externUID == "" {
return nil, status.Errorf(codes.InvalidArgument, "identity provider returned an empty subject identifier")
}
user, err := s.getLinkedSSOUser(ctx, provider, externUID)
if err != nil {
@@ -52,55 +58,86 @@ func (s *APIV1Service) resolveSSOUser(ctx context.Context, currentUser *store.Us
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),
})
user, err = s.createSSOUser(ctx, userInfo, string(passwordHash), provider, externUID)
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
}
// createSSOUser prefers the mapped external identifier as the initial local
// username when it satisfies the local username rules and is not a reserved
// name. A database uniqueness conflict falls back to a generated UUID instead of
// linking the SSO identity to the existing same-named account. User and identity
// creation are committed atomically so a concurrent UUID fallback cannot win
// after another request has claimed the preferred username.
//
// tryUsername returns a non-nil user when the identity is resolved (either newly
// created or reconciled to a concurrent winner) and (nil, nil) when the username
// is already taken and the caller should retry with a different one.
func (s *APIV1Service) createSSOUser(
ctx context.Context,
userInfo *idp.IdentityProviderUserInfo,
passwordHash string,
provider string,
externUID string,
) (*store.User, error) {
tryUsername := func(username string) (*store.User, error) {
user, err := s.Store.CreateUserWithIdentity(ctx, &store.User{
Username: username,
Role: store.RoleUser,
Nickname: userInfo.DisplayName,
Email: userInfo.Email,
AvatarURL: userInfo.AvatarURL,
PasswordHash: passwordHash,
}, &store.UserIdentity{
Provider: provider,
ExternUID: externUID,
})
if err == nil {
return user, nil
}
if !isUniqueConstraintViolation(err) {
return nil, err
}
// A unique violation is either the (provider, extern_uid) linkage (a
// concurrent first login won — reconcile to its user) or the username (in
// use by another account — signal a retry with a fresh username).
return s.getLinkedSSOUser(ctx, provider, externUID)
}
// Only adopt the external identifier as the local username when it is a valid,
// non-reserved name; otherwise an attacker-influenceable identifier could
// squat a privileged or system handle. Reserved and invalid names fall back to
// an opaque UUID.
if err := validateWritableUsername(userInfo.Identifier); err == nil && !isReservedUsername(userInfo.Identifier) {
user, err := tryUsername(userInfo.Identifier)
if err != nil {
return nil, err
}
if user != nil {
return user, nil
}
}
for range ssoUsernameFallbackAttempts {
username, err := deriveSSOUsername()
if err != nil {
return nil, err
}
user, err := tryUsername(username)
if err != nil {
return nil, err
}
if user != nil {
return user, nil
}
}
return nil, errors.Errorf("exhausted %d UUID username attempts", ssoUsernameFallbackAttempts)
}
func (s *APIV1Service) resolveSSOIdentity(ctx context.Context, idpName, code, redirectURI, codeVerifier string) (*storepb.IdentityProvider, *idp.IdentityProviderUserInfo, error) {
idpUID, err := ExtractIdentityProviderUIDFromName(idpName)
if err != nil {
@@ -225,8 +262,8 @@ func (s *APIV1Service) bindSSOIdentityToUser(ctx context.Context, currentUser *s
// 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.
// user row can only violate UNIQUE(username)). Shared by the SSO create/link paths
// and CreateMemo's UID uniqueness check.
func isUniqueConstraintViolation(err error) bool {
if err == nil {
return false
+2 -6
View File
@@ -5,7 +5,6 @@ import (
stderrors "errors"
"fmt"
"log/slog"
"strings"
"time"
"github.com/pkg/errors"
@@ -114,11 +113,8 @@ func (s *APIV1Service) CreateMemo(ctx context.Context, request *v1pb.CreateMemoR
memo, err := s.Store.CreateMemo(ctx, create)
if err != nil {
// Check for unique constraint violation (AIP-133 compliance)
errMsg := err.Error()
if strings.Contains(errMsg, "UNIQUE constraint failed") ||
strings.Contains(errMsg, "duplicate key") ||
strings.Contains(errMsg, "Duplicate entry") {
// Check for unique constraint violation (AIP-133 compliance).
if isUniqueConstraintViolation(err) {
return nil, status.Errorf(codes.AlreadyExists, "memo with ID %q already exists", memoUID)
}
return nil, err
+28 -3
View File
@@ -1,16 +1,41 @@
package v1
import (
"strings"
"github.com/pkg/errors"
"github.com/usememos/memos/internal/util"
)
const ssoUsernameFallbackAttempts = 5
// reservedUsernames are names that must never be auto-assigned from an external
// identity provider. A user-influenceable identifier (e.g. an OIDC
// preferred_username) claiming one of these on first login would squat a
// privileged or system-suggestive handle, so it falls back to a UUID instead.
var reservedUsernames = map[string]struct{}{
"admin": {},
"administrator": {},
"api": {},
"memos": {},
"root": {},
"support": {},
"system": {},
}
// isReservedUsername reports whether username is reserved. The comparison is
// case-insensitive because some backends fold case on the username unique index.
func isReservedUsername(username string) bool {
_, ok := reservedUsernames[strings.ToLower(strings.TrimSpace(username))]
return ok
}
// 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.
// UUID usernames are the fallback when the IdP identifier cannot safely be used
// as the local username, such as when it is invalid or already belongs to
// another local account.
func deriveSSOUsername() (string, error) {
username := util.GenUUID()
if err := validateWritableUsername(username); err != nil {
@@ -0,0 +1,383 @@
package test
import (
"context"
"strings"
"sync"
"testing"
"github.com/google/uuid"
"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 TestSSOSignInUsesValidIdentifierAsUsername(t *testing.T) {
tests := []struct {
name string
identifier string
}{
{name: "single character", identifier: "a"},
{name: "common username", identifier: "alice"},
{name: "uppercase and hyphens", identifier: "Alice-01"},
{name: "maximum length", identifier: strings.Repeat("a", 36)},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
ts := NewTestService(t)
defer ts.Cleanup()
ctx := context.Background()
mockIDP := newMockOAuthServer(t, "valid-code", "valid-token", map[string]any{
"sub": test.identifier,
"name": "Different Display Name",
"email": "alice@example.com",
})
defer mockIDP.Close()
idpName := createTestingOAuthIdentityProvider(ctx, t, ts, mockIDP.URL, "valid-identifier")
response, err := signInWithTestingSSO(ctx, ts, idpName, "valid-code")
require.NoError(t, err)
require.Equal(t, test.identifier, response.User.Username)
assertSingleSSOLink(ctx, t, ts, "valid-identifier", test.identifier, response.User.Username)
})
}
}
func TestSSOSignInFallsBackForInvalidIdentifier(t *testing.T) {
tests := []struct {
name string
identifier string
}{
{name: "numeric", identifier: "12345"},
{name: "email", identifier: "alice@example.com"},
{name: "underscore", identifier: "alice_example"},
{name: "leading hyphen", identifier: "-alice"},
{name: "trailing hyphen", identifier: "alice-"},
{name: "surrounding whitespace", identifier: " alice "},
{name: "too long", identifier: strings.Repeat("a", 37)},
{name: "non ASCII", identifier: "爱丽丝"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
ts := NewTestService(t)
defer ts.Cleanup()
ctx := context.Background()
mockIDP := newMockOAuthServer(t, "invalid-code", "invalid-token", map[string]any{
"sub": test.identifier,
"name": "Alice Example",
})
defer mockIDP.Close()
idpName := createTestingOAuthIdentityProvider(ctx, t, ts, mockIDP.URL, "invalid-identifier")
response, err := signInWithTestingSSO(ctx, ts, idpName, "invalid-code")
require.NoError(t, err)
require.NotEqual(t, test.identifier, response.User.Username)
_, err = uuid.Parse(response.User.Username)
require.NoError(t, err, "fallback username must be a UUID")
assertSingleSSOLink(ctx, t, ts, "invalid-identifier", test.identifier, response.User.Username)
})
}
}
func TestSSOSignInDoesNotTakeOverExistingUsername(t *testing.T) {
ts := NewTestService(t)
defer ts.Cleanup()
ctx := context.Background()
existingUser, err := ts.CreateRegularUser(ctx, "alice")
require.NoError(t, err)
mockIDP := newMockOAuthServer(t, "collision-code", "collision-token", map[string]any{
"sub": "alice",
"name": "SSO Alice",
"email": "sso-alice@example.com",
})
defer mockIDP.Close()
idpName := createTestingOAuthIdentityProvider(ctx, t, ts, mockIDP.URL, "username-collision")
response, err := signInWithTestingSSO(ctx, ts, idpName, "collision-code")
require.NoError(t, err)
require.NotEqual(t, existingUser.Username, response.User.Username)
_, err = uuid.Parse(response.User.Username)
require.NoError(t, err)
repeated, err := signInWithTestingSSO(ctx, ts, idpName, "collision-code")
require.NoError(t, err)
require.Equal(t, response.User.Name, repeated.User.Name)
stillExisting, err := ts.Store.GetUser(ctx, &store.FindUser{ID: &existingUser.ID})
require.NoError(t, err)
require.Equal(t, "alice", stillExisting.Username)
users, err := ts.Store.ListUsers(ctx, &store.FindUser{})
require.NoError(t, err)
require.Len(t, users, 2)
assertSingleSSOLink(ctx, t, ts, "username-collision", "alice", response.User.Username)
}
func TestSSOSignInDoesNotAdoptReservedUsername(t *testing.T) {
for _, identifier := range []string{"admin", "Admin", "support", "root"} {
t.Run(identifier, func(t *testing.T) {
ts := NewTestService(t)
defer ts.Cleanup()
ctx := context.Background()
mockIDP := newMockOAuthServer(t, "reserved-code", "reserved-token", map[string]any{"sub": identifier})
defer mockIDP.Close()
idpName := createTestingOAuthIdentityProvider(ctx, t, ts, mockIDP.URL, "reserved-provider")
response, err := signInWithTestingSSO(ctx, ts, idpName, "reserved-code")
require.NoError(t, err)
require.NotEqual(t, identifier, response.User.Username)
_, err = uuid.Parse(response.User.Username)
require.NoError(t, err, "reserved identifier must fall back to a UUID")
})
}
}
func TestSSOSignInReusesLinkedUser(t *testing.T) {
tests := []struct {
name string
identifier string
}{
{name: "preferred username", identifier: "alice"},
{name: "UUID fallback", identifier: "alice@example.com"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
ts := NewTestService(t)
defer ts.Cleanup()
ctx := context.Background()
mockIDP := newMockOAuthServer(t, "repeat-code", "repeat-token", map[string]any{
"sub": test.identifier,
"name": "Alice Example",
})
defer mockIDP.Close()
idpName := createTestingOAuthIdentityProvider(ctx, t, ts, mockIDP.URL, "repeat-provider")
first, err := signInWithTestingSSO(ctx, ts, idpName, "repeat-code")
require.NoError(t, err)
second, err := signInWithTestingSSO(ctx, ts, idpName, "repeat-code")
require.NoError(t, err)
require.Equal(t, first.User.Name, second.User.Name)
require.Equal(t, first.User.Username, second.User.Username)
users, err := ts.Store.ListUsers(ctx, &store.FindUser{})
require.NoError(t, err)
require.Len(t, users, 1)
assertSingleSSOLink(ctx, t, ts, "repeat-provider", test.identifier, first.User.Username)
})
}
}
func TestSSOSignInScopesSameIdentifierByProvider(t *testing.T) {
ts := NewTestService(t)
defer ts.Cleanup()
ctx := context.Background()
firstMockIDP := newMockOAuthServer(t, "first-code", "first-token", map[string]any{"sub": "alice"})
defer firstMockIDP.Close()
secondMockIDP := newMockOAuthServer(t, "second-code", "second-token", map[string]any{"sub": "alice"})
defer secondMockIDP.Close()
firstIDPName := createTestingOAuthIdentityProvider(ctx, t, ts, firstMockIDP.URL, "provider-one")
secondIDPName := createTestingOAuthIdentityProvider(ctx, t, ts, secondMockIDP.URL, "provider-two")
first, err := signInWithTestingSSO(ctx, ts, firstIDPName, "first-code")
require.NoError(t, err)
second, err := signInWithTestingSSO(ctx, ts, secondIDPName, "second-code")
require.NoError(t, err)
require.Equal(t, "alice", first.User.Username)
require.NotEqual(t, first.User.Name, second.User.Name)
_, err = uuid.Parse(second.User.Username)
require.NoError(t, err)
users, err := ts.Store.ListUsers(ctx, &store.FindUser{})
require.NoError(t, err)
require.Len(t, users, 2)
identities, err := ts.Store.ListUserIdentities(ctx, &store.FindUserIdentity{ExternUID: ptr("alice")})
require.NoError(t, err)
require.Len(t, identities, 2)
}
func TestConcurrentSSOFirstSignInConvergesOnOneUser(t *testing.T) {
ts := NewTestService(t)
defer ts.Cleanup()
ctx := context.Background()
mockIDP := newMockOAuthServer(t, "concurrent-code", "concurrent-token", map[string]any{
"sub": "alice",
"name": "Alice Example",
})
defer mockIDP.Close()
idpName := createTestingOAuthIdentityProvider(ctx, t, ts, mockIDP.URL, "concurrent-provider")
const signInCount = 8
start := make(chan struct{})
results := make(chan *v1pb.SignInResponse, signInCount)
errs := make(chan error, signInCount)
var waitGroup sync.WaitGroup
for range signInCount {
waitGroup.Add(1)
go func() {
defer waitGroup.Done()
<-start
response, err := signInWithTestingSSO(ctx, ts, idpName, "concurrent-code")
results <- response
errs <- err
}()
}
close(start)
waitGroup.Wait()
close(results)
close(errs)
for err := range errs {
require.NoError(t, err)
}
var userName string
for response := range results {
require.NotNil(t, response)
if userName == "" {
userName = response.User.Name
} else {
require.Equal(t, userName, response.User.Name)
}
}
users, err := ts.Store.ListUsers(ctx, &store.FindUser{})
require.NoError(t, err)
require.Len(t, users, 1)
assertSingleSSOLink(ctx, t, ts, "concurrent-provider", "alice", "alice")
}
func TestSSOSignInRejectsEmptyIdentifierWithoutCreatingUser(t *testing.T) {
ts := NewTestService(t)
defer ts.Cleanup()
ctx := context.Background()
mockIDP := newMockOAuthServer(t, "empty-code", "empty-token", map[string]any{"sub": ""})
defer mockIDP.Close()
idpName := createTestingOAuthIdentityProvider(ctx, t, ts, mockIDP.URL, "empty-identifier")
// An empty subject must never provision an account. The OAuth2 layer rejects it
// today, and resolveSSOUser guards it independently, so assert the invariant
// that matters — no user or identity is created — rather than which layer's
// error code surfaces.
_, err := signInWithTestingSSO(ctx, ts, idpName, "empty-code")
require.Error(t, err)
require.NotEqual(t, codes.OK, status.Code(err))
users, listErr := ts.Store.ListUsers(ctx, &store.FindUser{})
require.NoError(t, listErr)
require.Empty(t, users)
identities, listErr := ts.Store.ListUserIdentities(ctx, &store.FindUserIdentity{})
require.NoError(t, listErr)
require.Empty(t, identities)
}
func TestSSOSignInHonorsRegistrationGate(t *testing.T) {
ts := NewTestService(t)
defer ts.Cleanup()
ctx := context.Background()
_, err := ts.Store.UpsertInstanceSetting(ctx, &storepb.InstanceSetting{
Key: storepb.InstanceSettingKey_GENERAL,
Value: &storepb.InstanceSetting_GeneralSetting{
GeneralSetting: &storepb.InstanceGeneralSetting{DisallowUserRegistration: true},
},
})
require.NoError(t, err)
mockIDP := newMockOAuthServer(t, "blocked-code", "blocked-token", map[string]any{"sub": "alice"})
defer mockIDP.Close()
idpName := createTestingOAuthIdentityProvider(ctx, t, ts, mockIDP.URL, "blocked-provider")
_, err = signInWithTestingSSO(ctx, ts, idpName, "blocked-code")
require.Error(t, err)
require.Equal(t, codes.PermissionDenied, status.Code(err))
users, listErr := ts.Store.ListUsers(ctx, &store.FindUser{})
require.NoError(t, listErr)
require.Empty(t, users)
}
func TestSSOSignInAllowsLinkedUserWhenRegistrationDisabled(t *testing.T) {
ts := NewTestService(t)
defer ts.Cleanup()
ctx := context.Background()
mockIDP := newMockOAuthServer(t, "linked-code", "linked-token", map[string]any{"sub": "alice"})
defer mockIDP.Close()
idpName := createTestingOAuthIdentityProvider(ctx, t, ts, mockIDP.URL, "linked-provider")
first, err := signInWithTestingSSO(ctx, ts, idpName, "linked-code")
require.NoError(t, err)
_, err = ts.Store.UpsertInstanceSetting(ctx, &storepb.InstanceSetting{
Key: storepb.InstanceSettingKey_GENERAL,
Value: &storepb.InstanceSetting_GeneralSetting{
GeneralSetting: &storepb.InstanceGeneralSetting{DisallowUserRegistration: true},
},
})
require.NoError(t, err)
second, err := signInWithTestingSSO(ctx, ts, idpName, "linked-code")
require.NoError(t, err)
require.Equal(t, first.User.Name, second.User.Name)
users, err := ts.Store.ListUsers(ctx, &store.FindUser{})
require.NoError(t, err)
require.Len(t, users, 1)
}
func signInWithTestingSSO(ctx context.Context, ts *TestService, idpName, code string) (*v1pb.SignInResponse, error) {
return ts.Service.SignIn(apiv1.WithHeaderCarrier(ctx), &v1pb.SignInRequest{
Credentials: &v1pb.SignInRequest_SsoCredentials{
SsoCredentials: &v1pb.SignInRequest_SSOCredentials{
IdpName: idpName,
Code: code,
RedirectUri: "http://localhost:8080/auth/callback",
},
},
})
}
func assertSingleSSOLink(ctx context.Context, t *testing.T, ts *TestService, provider, externUID, username string) {
t.Helper()
identity, err := ts.Store.GetUserIdentity(ctx, &store.FindUserIdentity{
Provider: &provider,
ExternUID: &externUID,
})
require.NoError(t, err)
require.NotNil(t, identity)
user, err := ts.Store.GetUser(ctx, &store.FindUser{ID: &identity.UserID})
require.NoError(t, err)
require.NotNil(t, user)
require.Equal(t, username, user.Username)
identities, err := ts.Store.ListUserIdentities(ctx, &store.FindUserIdentity{
Provider: &provider,
ExternUID: &externUID,
})
require.NoError(t, err)
require.Len(t, identities, 1)
}
func ptr[T any](value T) *T {
return &value
}
+1 -6
View File
@@ -11,12 +11,7 @@ import (
)
func (d *DB) CreateUser(ctx context.Context, create *store.User) (*store.User, error) {
fields := []string{"`username`", "`role`", "`email`", "`nickname`", "`password_hash`", "`avatar_url`"}
placeholder := []string{"?", "?", "?", "?", "?", "?"}
args := []any{create.Username, create.Role, create.Email, create.Nickname, create.PasswordHash, create.AvatarURL}
stmt := "INSERT INTO user (" + strings.Join(fields, ", ") + ") VALUES (" + strings.Join(placeholder, ", ") + ")"
result, err := d.db.ExecContext(ctx, stmt, args...)
result, err := insertUser(ctx, d.db, create)
if err != nil {
return nil, err
}
+62 -2
View File
@@ -2,6 +2,7 @@ package mysql
import (
"context"
"database/sql"
"strings"
"github.com/pkg/errors"
@@ -9,9 +10,24 @@ import (
"github.com/usememos/memos/store"
)
func (d *DB) CreateUserIdentity(ctx context.Context, create *store.UserIdentity) (*store.UserIdentity, error) {
// execer is satisfied by both *sql.DB and *sql.Tx so the INSERT statements can be
// shared between the standalone and transactional creation paths.
type execer interface {
ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error)
}
func insertUser(ctx context.Context, e execer, create *store.User) (sql.Result, error) {
stmt := "INSERT INTO user (`username`, `role`, `email`, `nickname`, `password_hash`, `avatar_url`) VALUES (?, ?, ?, ?, ?, ?)"
return e.ExecContext(ctx, stmt, create.Username, create.Role, create.Email, create.Nickname, create.PasswordHash, create.AvatarURL)
}
func insertUserIdentity(ctx context.Context, e execer, create *store.UserIdentity) (sql.Result, error) {
stmt := "INSERT INTO `user_identity` (`user_id`, `provider`, `extern_uid`) VALUES (?, ?, ?)"
result, err := d.db.ExecContext(ctx, stmt, create.UserID, create.Provider, create.ExternUID)
return e.ExecContext(ctx, stmt, create.UserID, create.Provider, create.ExternUID)
}
func (d *DB) CreateUserIdentity(ctx context.Context, create *store.UserIdentity) (*store.UserIdentity, error) {
result, err := insertUserIdentity(ctx, d.db, create)
if err != nil {
return nil, err
}
@@ -31,6 +47,50 @@ func (d *DB) CreateUserIdentity(ctx context.Context, create *store.UserIdentity)
return list[0], nil
}
func (d *DB) CreateUserWithIdentity(ctx context.Context, createUser *store.User, createIdentity *store.UserIdentity) (*store.User, error) {
tx, err := d.db.BeginTx(ctx, nil)
if err != nil {
return nil, errors.Wrap(err, "failed to begin user identity transaction")
}
defer func() {
_ = tx.Rollback()
}()
userResult, err := insertUser(ctx, tx, createUser)
if err != nil {
return nil, errors.Wrap(err, "failed to create user")
}
rawUserID, err := userResult.LastInsertId()
if err != nil {
return nil, errors.Wrap(err, "failed to read created user ID")
}
createUser.ID = int32(rawUserID)
// RETURNING is unavailable on MySQL, so read back the DB-populated columns the
// user cache needs within the same transaction.
if err := tx.QueryRowContext(
ctx,
"SELECT `description`, UNIX_TIMESTAMP(`created_ts`), UNIX_TIMESTAMP(`updated_ts`), `row_status` FROM `user` WHERE `id` = ?",
createUser.ID,
).Scan(
&createUser.Description,
&createUser.CreatedTs,
&createUser.UpdatedTs,
&createUser.RowStatus,
); err != nil {
return nil, errors.Wrap(err, "failed to read created user")
}
createIdentity.UserID = createUser.ID
if _, err := insertUserIdentity(ctx, tx, createIdentity); err != nil {
return nil, errors.Wrap(err, "failed to create user identity")
}
if err := tx.Commit(); err != nil {
return nil, errors.Wrap(err, "failed to commit user identity transaction")
}
return createUser, nil
}
func (d *DB) ListUserIdentities(ctx context.Context, find *store.FindUserIdentity) ([]*store.UserIdentity, error) {
where, args := []string{"1 = 1"}, []any{}
+1 -11
View File
@@ -11,19 +11,9 @@ import (
)
func (d *DB) CreateUser(ctx context.Context, create *store.User) (*store.User, error) {
fields := []string{"username", "role", "email", "nickname", "password_hash", "avatar_url"}
args := []any{create.Username, create.Role, create.Email, create.Nickname, create.PasswordHash, create.AvatarURL}
stmt := "INSERT INTO \"user\" (" + strings.Join(fields, ", ") + ") VALUES (" + placeholders(len(args)) + ") RETURNING id, description, created_ts, updated_ts, row_status"
if err := d.db.QueryRowContext(ctx, stmt, args...).Scan(
&create.ID,
&create.Description,
&create.CreatedTs,
&create.UpdatedTs,
&create.RowStatus,
); err != nil {
if err := insertUser(ctx, d.db, create); err != nil {
return nil, err
}
return create, nil
}
+60 -3
View File
@@ -2,23 +2,80 @@ package postgres
import (
"context"
"database/sql"
"strings"
"github.com/pkg/errors"
"github.com/usememos/memos/store"
)
func (d *DB) CreateUserIdentity(ctx context.Context, create *store.UserIdentity) (*store.UserIdentity, error) {
// rowQuerier is satisfied by both *sql.DB and *sql.Tx so the insert statements
// can be shared between the standalone and transactional creation paths.
type rowQuerier interface {
QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row
}
func insertUser(ctx context.Context, q rowQuerier, create *store.User) error {
stmt := "INSERT INTO \"user\" (username, role, email, nickname, password_hash, avatar_url) VALUES (" + placeholders(6) + ") RETURNING id, description, created_ts, updated_ts, row_status"
return q.QueryRowContext(
ctx,
stmt,
create.Username,
create.Role,
create.Email,
create.Nickname,
create.PasswordHash,
create.AvatarURL,
).Scan(
&create.ID,
&create.Description,
&create.CreatedTs,
&create.UpdatedTs,
&create.RowStatus,
)
}
func insertUserIdentity(ctx context.Context, q rowQuerier, create *store.UserIdentity) error {
stmt := "INSERT INTO user_identity (user_id, provider, extern_uid) VALUES (" + placeholders(3) + ") RETURNING id, created_ts, updated_ts"
if err := d.db.QueryRowContext(ctx, stmt, create.UserID, create.Provider, create.ExternUID).Scan(
return q.QueryRowContext(ctx, stmt, create.UserID, create.Provider, create.ExternUID).Scan(
&create.ID,
&create.CreatedTs,
&create.UpdatedTs,
); err != nil {
)
}
func (d *DB) CreateUserIdentity(ctx context.Context, create *store.UserIdentity) (*store.UserIdentity, error) {
if err := insertUserIdentity(ctx, d.db, create); err != nil {
return nil, err
}
return create, nil
}
func (d *DB) CreateUserWithIdentity(ctx context.Context, createUser *store.User, createIdentity *store.UserIdentity) (*store.User, error) {
tx, err := d.db.BeginTx(ctx, nil)
if err != nil {
return nil, errors.Wrap(err, "failed to begin user identity transaction")
}
defer func() {
_ = tx.Rollback()
}()
if err := insertUser(ctx, tx, createUser); err != nil {
return nil, errors.Wrap(err, "failed to create user")
}
createIdentity.UserID = createUser.ID
if err := insertUserIdentity(ctx, tx, createIdentity); err != nil {
return nil, errors.Wrap(err, "failed to create user identity")
}
if err := tx.Commit(); err != nil {
return nil, errors.Wrap(err, "failed to commit user identity transaction")
}
return createUser, nil
}
func (d *DB) ListUserIdentities(ctx context.Context, find *store.FindUserIdentity) ([]*store.UserIdentity, error) {
where, args := []string{"1 = 1"}, []any{}
+1 -12
View File
@@ -11,20 +11,9 @@ import (
)
func (d *DB) CreateUser(ctx context.Context, create *store.User) (*store.User, error) {
fields := []string{"`username`", "`role`", "`email`", "`nickname`", "`password_hash`, `avatar_url`"}
placeholder := []string{"?", "?", "?", "?", "?", "?"}
args := []any{create.Username, create.Role, create.Email, create.Nickname, create.PasswordHash, create.AvatarURL}
stmt := "INSERT INTO user (" + strings.Join(fields, ", ") + ") VALUES (" + strings.Join(placeholder, ", ") + ") RETURNING id, description, created_ts, updated_ts, row_status"
if err := d.db.QueryRowContext(ctx, stmt, args...).Scan(
&create.ID,
&create.Description,
&create.CreatedTs,
&create.UpdatedTs,
&create.RowStatus,
); err != nil {
if err := insertUser(ctx, d.db, create); err != nil {
return nil, err
}
return create, nil
}
+60 -3
View File
@@ -2,23 +2,80 @@ package sqlite
import (
"context"
"database/sql"
"strings"
"github.com/pkg/errors"
"github.com/usememos/memos/store"
)
func (d *DB) CreateUserIdentity(ctx context.Context, create *store.UserIdentity) (*store.UserIdentity, error) {
// rowQuerier is satisfied by both *sql.DB and *sql.Tx so the insert statements
// can be shared between the standalone and transactional creation paths.
type rowQuerier interface {
QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row
}
func insertUser(ctx context.Context, q rowQuerier, create *store.User) error {
stmt := "INSERT INTO user (`username`, `role`, `email`, `nickname`, `password_hash`, `avatar_url`) VALUES (?, ?, ?, ?, ?, ?) RETURNING id, description, created_ts, updated_ts, row_status"
return q.QueryRowContext(
ctx,
stmt,
create.Username,
create.Role,
create.Email,
create.Nickname,
create.PasswordHash,
create.AvatarURL,
).Scan(
&create.ID,
&create.Description,
&create.CreatedTs,
&create.UpdatedTs,
&create.RowStatus,
)
}
func insertUserIdentity(ctx context.Context, q rowQuerier, create *store.UserIdentity) error {
stmt := "INSERT INTO `user_identity` (`user_id`, `provider`, `extern_uid`) VALUES (?, ?, ?) RETURNING `id`, `created_ts`, `updated_ts`"
if err := d.db.QueryRowContext(ctx, stmt, create.UserID, create.Provider, create.ExternUID).Scan(
return q.QueryRowContext(ctx, stmt, create.UserID, create.Provider, create.ExternUID).Scan(
&create.ID,
&create.CreatedTs,
&create.UpdatedTs,
); err != nil {
)
}
func (d *DB) CreateUserIdentity(ctx context.Context, create *store.UserIdentity) (*store.UserIdentity, error) {
if err := insertUserIdentity(ctx, d.db, create); err != nil {
return nil, err
}
return create, nil
}
func (d *DB) CreateUserWithIdentity(ctx context.Context, createUser *store.User, createIdentity *store.UserIdentity) (*store.User, error) {
tx, err := d.db.BeginTx(ctx, nil)
if err != nil {
return nil, errors.Wrap(err, "failed to begin user identity transaction")
}
defer func() {
_ = tx.Rollback()
}()
if err := insertUser(ctx, tx, createUser); err != nil {
return nil, errors.Wrap(err, "failed to create user")
}
createIdentity.UserID = createUser.ID
if err := insertUserIdentity(ctx, tx, createIdentity); err != nil {
return nil, errors.Wrap(err, "failed to create user identity")
}
if err := tx.Commit(); err != nil {
return nil, errors.Wrap(err, "failed to commit user identity transaction")
}
return createUser, nil
}
func (d *DB) ListUserIdentities(ctx context.Context, find *store.FindUserIdentity) ([]*store.UserIdentity, error) {
where, args := []string{"1 = 1"}, []any{}
+1
View File
@@ -81,6 +81,7 @@ type Driver interface {
// UserIdentity model related methods.
CreateUserIdentity(ctx context.Context, create *UserIdentity) (*UserIdentity, error)
CreateUserWithIdentity(ctx context.Context, createUser *User, createIdentity *UserIdentity) (*User, error)
ListUserIdentities(ctx context.Context, find *FindUserIdentity) ([]*UserIdentity, error)
DeleteUserIdentities(ctx context.Context, delete *DeleteUserIdentity) error
}
+74
View File
@@ -161,6 +161,80 @@ func TestUserIdentitySameUserSameProviderConflicts(t *testing.T) {
require.Error(t, err)
}
func TestCreateUserWithIdentityIsAtomic(t *testing.T) {
t.Parallel()
ctx := context.Background()
ts := NewTestingStore(ctx, t)
defer ts.Close()
t.Run("creates both records", func(t *testing.T) {
provider := "atomic-provider"
externUID := "atomic-subject"
user, err := ts.CreateUserWithIdentity(ctx, &store.User{
Username: "atomic-user",
Role: store.RoleUser,
}, &store.UserIdentity{
Provider: provider,
ExternUID: externUID,
})
require.NoError(t, err)
require.NotZero(t, user.ID)
storedIdentity, err := ts.GetUserIdentity(ctx, &store.FindUserIdentity{Provider: &provider, ExternUID: &externUID})
require.NoError(t, err)
require.NotNil(t, storedIdentity)
require.Equal(t, user.ID, storedIdentity.UserID)
})
t.Run("rolls back user when identity conflicts", func(t *testing.T) {
owner, err := createTestingUserWithRole(ctx, ts, "identity-owner", store.RoleUser)
require.NoError(t, err)
_, err = ts.CreateUserIdentity(ctx, &store.UserIdentity{
UserID: owner.ID,
Provider: "conflict-provider",
ExternUID: "conflict-subject",
})
require.NoError(t, err)
_, err = ts.CreateUserWithIdentity(ctx, &store.User{
Username: "rolled-back-user",
Role: store.RoleUser,
}, &store.UserIdentity{
Provider: "conflict-provider",
ExternUID: "conflict-subject",
})
require.Error(t, err)
username := "rolled-back-user"
user, err := ts.GetUser(ctx, &store.FindUser{Username: &username})
require.NoError(t, err)
require.Nil(t, user)
})
t.Run("rolls back identity when username conflicts", func(t *testing.T) {
_, err := createTestingUserWithRole(ctx, ts, "taken-username", store.RoleUser)
require.NoError(t, err)
_, err = ts.CreateUserWithIdentity(ctx, &store.User{
Username: "taken-username",
Role: store.RoleUser,
}, &store.UserIdentity{
Provider: "unused-provider",
ExternUID: "unused-subject",
})
require.Error(t, err)
provider := "unused-provider"
externUID := "unused-subject"
identity, err := ts.GetUserIdentity(ctx, &store.FindUserIdentity{
Provider: &provider,
ExternUID: &externUID,
})
require.NoError(t, err)
require.Nil(t, identity)
})
}
func TestUserIdentityDeleteByUserAndProvider(t *testing.T) {
t.Parallel()
ctx := context.Background()
+11
View File
@@ -36,6 +36,17 @@ func (s *Store) CreateUserIdentity(ctx context.Context, create *UserIdentity) (*
return s.driver.CreateUserIdentity(ctx, create)
}
// CreateUserWithIdentity atomically creates a local user and its external identity
// linkage, returning the created user.
func (s *Store) CreateUserWithIdentity(ctx context.Context, createUser *User, createIdentity *UserIdentity) (*User, error) {
user, err := s.driver.CreateUserWithIdentity(ctx, createUser, createIdentity)
if err != nil {
return nil, err
}
s.userCache.Set(ctx, userCacheKey(user.ID), user)
return user, nil
}
// ListUserIdentities returns all linkage records matching the filter.
func (s *Store) ListUserIdentities(ctx context.Context, find *FindUserIdentity) ([]*UserIdentity, error) {
return s.driver.ListUserIdentities(ctx, find)