diff --git a/backend/pkg/server/services/auth.go b/backend/pkg/server/services/auth.go
index 682eed5c..81faea07 100644
--- a/backend/pkg/server/services/auth.go
+++ b/backend/pkg/server/services/auth.go
@@ -450,7 +450,6 @@ func (s *AuthService) AuthLogout(c *gin.Context) {
func (s *AuthService) authLoginCallback(c *gin.Context, stateData map[string]string, code string) {
var (
privs []string
- role models.Role
user models.User
)
@@ -504,24 +503,7 @@ func (s *AuthService) authLoginCallback(c *gin.Context, stateData map[string]str
return
}
- err = s.db.Take(&role, "id = ?", models.RoleUser).Error
- if err != nil {
- logger.FromContext(c).WithError(err).Errorf("error getting user role '%d'", models.RoleUser)
- response.Error(c, response.ErrAuthInvalidServiceData, err)
- return
- }
-
- err = s.db.Table("privileges").
- Where("role_id = ?", models.RoleUser).
- Pluck("name", &privs).Error
- if err != nil {
- logger.FromContext(c).WithError(err).Errorf("error getting user privileges list '%s'", user.Hash)
- response.Error(c, response.ErrAuthInvalidServiceData, err)
- return
- }
-
- filterQuery := "mail = ? AND type = ?"
- if err = s.db.Take(&user, filterQuery, email, models.UserTypeOAuth).Error; err != nil {
+ if err = s.db.Take(&user, "mail = ?", email).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
user = models.User{
Hash: rdb.MakeUserHash(email),
@@ -542,23 +524,30 @@ func (s *AuthService) authLoginCallback(c *gin.Context, stateData map[string]str
if err = tx.Create(&user).Error; err != nil {
tx.Rollback()
- logger.FromContext(c).WithError(err).Errorf("error creating user")
- response.Error(c, response.ErrInternal, err)
- return
- }
+ if !isUniqueViolation(err) {
+ logger.FromContext(c).WithError(err).Errorf("error creating user")
+ response.Error(c, response.ErrInternal, err)
+ return
+ }
+ if err = s.db.Take(&user, "mail = ?", email).Error; err != nil {
+ logger.FromContext(c).WithError(err).Errorf("error loading concurrently created user '%s'", email)
+ response.Error(c, response.ErrInternal, err)
+ return
+ }
+ } else {
+ preferences := models.NewUserPreferences(user.ID)
+ if err = tx.Create(preferences).Error; err != nil {
+ tx.Rollback()
+ logger.FromContext(c).WithError(err).Errorf("error creating user preferences")
+ response.Error(c, response.ErrInternal, err)
+ return
+ }
- preferences := models.NewUserPreferences(user.ID)
- if err = tx.Create(preferences).Error; err != nil {
- tx.Rollback()
- logger.FromContext(c).WithError(err).Errorf("error creating user preferences")
- response.Error(c, response.ErrInternal, err)
- return
- }
-
- if err = tx.Commit().Error; err != nil {
- logger.FromContext(c).WithError(err).Errorf("error committing transaction")
- response.Error(c, response.ErrInternal, err)
- return
+ if err = tx.Commit().Error; err != nil {
+ logger.FromContext(c).WithError(err).Errorf("error committing transaction")
+ response.Error(c, response.ErrInternal, err)
+ return
+ }
}
} else {
logger.FromContext(c).WithError(err).Errorf("error searching user by email '%s'", email)
@@ -581,6 +570,15 @@ func (s *AuthService) authLoginCallback(c *gin.Context, stateData map[string]str
return
}
+ err = s.db.Table("privileges").
+ Where("role_id = ?", user.RoleID).
+ Pluck("name", &privs).Error
+ if err != nil {
+ logger.FromContext(c).WithError(err).Errorf("error getting user privileges list '%s'", user.Hash)
+ response.Error(c, response.ErrAuthInvalidServiceData, err)
+ return
+ }
+
expires := s.cfg.SessionTimeout
gtm := time.Now().Unix()
exp := time.Now().Add(time.Duration(expires) * time.Second).Unix()
diff --git a/backend/pkg/server/services/auth_test.go b/backend/pkg/server/services/auth_test.go
new file mode 100644
index 00000000..384c26f6
--- /dev/null
+++ b/backend/pkg/server/services/auth_test.go
@@ -0,0 +1,210 @@
+package services
+
+import (
+ "context"
+ "net/http"
+ "net/http/httptest"
+ "testing"
+ "time"
+
+ "github.com/gin-contrib/sessions"
+ "github.com/gin-contrib/sessions/cookie"
+ "github.com/gin-gonic/gin"
+ "github.com/jinzhu/gorm"
+ _ "github.com/jinzhu/gorm/dialects/sqlite"
+ "github.com/stretchr/testify/assert"
+ "github.com/stretchr/testify/require"
+ "golang.org/x/oauth2"
+
+ "pentagi/pkg/server/models"
+ "pentagi/pkg/server/oauth"
+)
+
+// fakeOAuthClient is an OAuthClient stub that returns a fixed, already-verified email
+// (the real google/github resolvers reject unverified addresses upstream).
+type fakeOAuthClient struct {
+ name string
+ email string
+}
+
+func (f *fakeOAuthClient) ProviderName() string { return f.name }
+
+func (f *fakeOAuthClient) ResolveEmail(context.Context, string, *oauth2.Token) (string, error) {
+ return f.email, nil
+}
+
+func (f *fakeOAuthClient) TokenSource(context.Context, *oauth2.Token) oauth2.TokenSource { return nil }
+
+func (f *fakeOAuthClient) Exchange(context.Context, string, ...oauth2.AuthCodeOption) (*oauth2.Token, error) {
+ return &oauth2.Token{AccessToken: "test-access-token", Expiry: time.Now().Add(time.Hour)}, nil
+}
+
+func (f *fakeOAuthClient) RefreshToken(context.Context, string) (*oauth2.Token, error) {
+ return nil, nil
+}
+
+func (f *fakeOAuthClient) AuthCodeURL(string, ...oauth2.AuthCodeOption) string { return "" }
+
+func newOAuthService(db *gorm.DB, email string) *AuthService {
+ return &AuthService{
+ cfg: AuthServiceConfig{BaseURL: "/", SessionTimeout: 3600},
+ db: db,
+ key: []byte("0123456789abcdef0123456789abcdef"),
+ oauth: map[string]oauth.OAuthClient{"github": &fakeOAuthClient{name: "github", email: email}},
+ }
+}
+
+func newCallbackContext(t *testing.T) (*gin.Context, *httptest.ResponseRecorder) {
+ t.Helper()
+ gin.SetMode(gin.TestMode)
+ w := httptest.NewRecorder()
+ c, _ := gin.CreateTestContext(w)
+
+ req := httptest.NewRequest(http.MethodGet, "/callback", nil)
+ req.AddCookie(&http.Cookie{Name: authNonceCookieName, Value: "test-nonce"})
+ c.Request = req
+
+ sessions.Sessions("pentagi", cookie.NewStore([]byte("test-secret")))(c)
+
+ return c, w
+}
+
+func countUsers(t *testing.T, db *gorm.DB) int {
+ t.Helper()
+ var count int
+ require.NoError(t, db.Model(&models.User{}).Count(&count).Error)
+ return count
+}
+
+// TestAuthLoginCallback_LinksExistingLocalAccount is the regression guard for the OAuth
+// squat DoS: a first OAuth login for an email already held by a LOCAL account must link
+// into that account (logging the verified owner in and recording the provider) rather than
+// taking the create branch, which would violate users_mail_unique and return 500.
+func TestAuthLoginCallback_LinksExistingLocalAccount(t *testing.T) {
+ db := setupTestDB(t)
+ defer db.Close()
+
+ require.NoError(t, db.Exec(
+ "INSERT INTO users (id, hash, type, mail, name, status, role_id, provider) VALUES (10, ?, 'local', 'victim@corp.com', 'Victim', 'active', 2, NULL)",
+ "1234567890abcdef1234567890abcdef",
+ ).Error)
+
+ before := countUsers(t, db)
+
+ svc := newOAuthService(db, "victim@corp.com")
+ c, w := newCallbackContext(t)
+ svc.authLoginCallback(c, map[string]string{"provider": "github"}, "test-code")
+
+ assert.Equal(t, http.StatusOK, w.Code)
+ assert.Equal(t, before, countUsers(t, db), "must link into the existing row, not create a duplicate")
+
+ var linked models.User
+ require.NoError(t, db.Where("mail = ?", "victim@corp.com").First(&linked).Error)
+ assert.Equal(t, uint64(10), linked.ID, "the existing local row is reused")
+ assert.Equal(t, models.UserTypeLocal, linked.Type, "linking keeps the password-login capability")
+ require.NotNil(t, linked.Provider)
+ assert.Equal(t, "github", *linked.Provider, "the provider is backfilled on link")
+
+ assert.Equal(t, uint64(10), sessions.Default(c).Get("uid"), "session is issued for the linked account")
+}
+
+func TestAuthLoginCallback_CreatesUserWhenEmailFree(t *testing.T) {
+ db := setupTestDB(t)
+ defer db.Close()
+
+ before := countUsers(t, db)
+
+ svc := newOAuthService(db, "newcomer@corp.com")
+ c, w := newCallbackContext(t)
+ svc.authLoginCallback(c, map[string]string{"provider": "github"}, "test-code")
+
+ assert.Equal(t, http.StatusOK, w.Code)
+ assert.Equal(t, before+1, countUsers(t, db))
+
+ var created models.User
+ require.NoError(t, db.Where("mail = ?", "newcomer@corp.com").First(&created).Error)
+ assert.Equal(t, models.UserTypeOAuth, created.Type)
+ require.NotNil(t, created.Provider)
+ assert.Equal(t, "github", *created.Provider)
+
+ var prefCount int
+ require.NoError(t, db.Table("user_preferences").Where("user_id = ?", created.ID).Count(&prefCount).Error)
+ assert.Equal(t, 1, prefCount, "preferences row is created alongside the new user")
+}
+
+// TestAuthLoginCallback_RejectsBlockedAccount guards the status gate: linking must not log in a
+// non-active account. Before the lookup was broadened, a blocked LOCAL row was invisible to the
+// callback (type filter) and a fresh active OAuth row was created instead — the gate now applies.
+func TestAuthLoginCallback_RejectsBlockedAccount(t *testing.T) {
+ db := setupTestDB(t)
+ defer db.Close()
+
+ require.NoError(t, db.Exec(
+ "INSERT INTO users (id, hash, type, mail, name, status, role_id) VALUES (11, ?, 'local', 'blocked@corp.com', 'Blocked', 'blocked', 2)",
+ "1234567890abcdef1234567890abcdef",
+ ).Error)
+ before := countUsers(t, db)
+
+ svc := newOAuthService(db, "blocked@corp.com")
+ c, w := newCallbackContext(t)
+ svc.authLoginCallback(c, map[string]string{"provider": "github"}, "test-code")
+
+ assert.Equal(t, http.StatusForbidden, w.Code, "a blocked account must not be logged in via OAuth")
+ assert.Nil(t, sessions.Default(c).Get("uid"), "no session is issued for a blocked account")
+ assert.Equal(t, before, countUsers(t, db), "no shadow account is created for a blocked email")
+}
+
+// TestAuthLoginCallback_LinkInheritsAccountRole pins that a linked session carries the privileges of
+// the account's actual role, not a hardcoded RoleUser set — rid and prm must agree. setupTestDB seeds
+// Admin (role 1) with users.create and User (role 2) without it.
+func TestAuthLoginCallback_LinkInheritsAccountRole(t *testing.T) {
+ db := setupTestDB(t)
+ defer db.Close()
+
+ require.NoError(t, db.Exec(
+ "INSERT INTO users (id, hash, type, mail, name, status, role_id) VALUES (12, ?, 'local', 'admin2@corp.com', 'Admin Two', 'active', 1)",
+ "1234567890abcdef1234567890abcdef",
+ ).Error)
+
+ svc := newOAuthService(db, "admin2@corp.com")
+ c, w := newCallbackContext(t)
+ svc.authLoginCallback(c, map[string]string{"provider": "github"}, "test-code")
+
+ require.Equal(t, http.StatusOK, w.Code)
+ sess := sessions.Default(c)
+ assert.Equal(t, uint64(1), sess.Get("rid"), "session role matches the linked account")
+ prm, _ := sess.Get("prm").([]string)
+ assert.Contains(t, prm, "users.create", "linked session carries the account role's privileges, not RoleUser")
+}
+
+// TestAuthLoginCallback_LinksOnConcurrentCreateConflict simulates the TOCTOU race where a second
+// first-login for the same new email commits between this request's lookup and its insert. The
+// one-shot callback hides the row from the first lookup, so the handler enters the create branch and
+// trips users_mail_unique; it must then re-fetch and link rather than return a 500.
+func TestAuthLoginCallback_LinksOnConcurrentCreateConflict(t *testing.T) {
+ db := setupTestDB(t)
+ defer db.Close()
+
+ require.NoError(t, db.Exec(
+ "INSERT INTO users (id, hash, type, mail, name, status, role_id) VALUES (20, ?, 'oauth', 'race@corp.com', 'Racer', 'active', 2)",
+ "1234567890abcdef1234567890abcdef",
+ ).Error)
+ before := countUsers(t, db)
+
+ hiddenOnce := false
+ db.Callback().Query().Before("gorm:query").Register("test:hide_first_user_query", func(scope *gorm.Scope) {
+ if !hiddenOnce && scope.TableName() == "users" {
+ hiddenOnce = true
+ scope.Err(gorm.ErrRecordNotFound)
+ }
+ })
+ defer db.Callback().Query().Remove("test:hide_first_user_query")
+
+ svc := newOAuthService(db, "race@corp.com")
+ c, w := newCallbackContext(t)
+ svc.authLoginCallback(c, map[string]string{"provider": "github"}, "test-code")
+
+ assert.Equal(t, http.StatusOK, w.Code, "a create conflict must relink, not 500")
+ assert.Equal(t, before, countUsers(t, db), "no duplicate is created on the conflict")
+ assert.Equal(t, uint64(20), sessions.Default(c).Get("uid"), "session is issued for the existing row")
+}
diff --git a/backend/pkg/server/services/resources.go b/backend/pkg/server/services/resources.go
index 8fae9956..b23a987e 100644
--- a/backend/pkg/server/services/resources.go
+++ b/backend/pkg/server/services/resources.go
@@ -2253,13 +2253,14 @@ func (s *ResourceService) publishResourcesDeleted(ctx context.Context, uid uint6
// ---- utility ---------------------------------------------------------------
-// isUniqueViolation returns true if err is a PostgreSQL unique constraint
-// violation (error code 23505).
+// isUniqueViolation returns true if err is a unique-constraint violation. Matches Postgres
+// ("duplicate key value violates unique constraint", SQLSTATE 23505) and SQLite ("UNIQUE
+// constraint failed") case-insensitively.
func isUniqueViolation(err error) bool {
if err == nil {
return false
}
- msg := err.Error()
+ msg := strings.ToLower(err.Error())
return strings.Contains(msg, "unique") ||
strings.Contains(msg, "duplicate") ||
strings.Contains(msg, "23505")
diff --git a/backend/pkg/server/services/users.go b/backend/pkg/server/services/users.go
index 860d4362..39a50b6a 100644
--- a/backend/pkg/server/services/users.go
+++ b/backend/pkg/server/services/users.go
@@ -248,11 +248,18 @@ func (s *UserService) ChangeEmailCurrentUser(c *gin.Context) {
return
}
+ // OAuth logins match accounts by email (authLoginCallback), so a new address unlinks the provider.
updates := map[string]any{
- "mail": form.Mail,
+ "mail": form.Mail,
+ "provider": nil,
}
if err = s.db.Model(&user).Scopes(scope).Updates(updates).Error; err != nil {
+ if isUniqueViolation(err) {
+ logger.FromContext(c).Warnf("email change rejected: address claimed concurrently")
+ response.Error(c, response.ErrChangeEmailCurrentUserEmailAlreadyExists, errors.New("email already exists"))
+ return
+ }
logger.FromContext(c).WithError(err).Errorf("error updating email for current user")
response.Error(c, response.ErrInternal, err)
return
diff --git a/backend/pkg/server/services/users_test.go b/backend/pkg/server/services/users_test.go
index 19c6427b..9de2803f 100644
--- a/backend/pkg/server/services/users_test.go
+++ b/backend/pkg/server/services/users_test.go
@@ -3,6 +3,7 @@ package services
import (
"bytes"
"encoding/json"
+ "errors"
"net/http"
"net/http/httptest"
"testing"
@@ -364,6 +365,7 @@ func TestChangeEmailCurrentUser(t *testing.T) {
err = db.Model(&models.User{}).Where("id = 1").Updates(map[string]interface{}{
"password": string(hashedPassword),
"hash": "11111111111111111111111111111111",
+ "provider": "github",
}).Error
require.NoError(t, err)
@@ -388,27 +390,28 @@ func TestChangeEmailCurrentUser(t *testing.T) {
err := db.Where("id = 1").First(&user).Error
require.NoError(t, err)
assert.Equal(t, "newemail@test.com", user.Mail)
+ assert.Nil(t, user.Provider, "email change clears the now-stale OAuth provider link")
},
},
{
name: "invalid password",
- requestBody: `{"current_password": "WrongPassword!", "mail": "another@test.com"}`,
- uid: 1,
- expectedCode: http.StatusForbidden,
+ requestBody: `{"current_password": "WrongPassword!", "mail": "another@test.com"}`,
+ uid: 1,
+ expectedCode: http.StatusForbidden,
errorContains: "invalid current password",
},
{
name: "email already exists",
- requestBody: `{"current_password": "SecurePass123!", "mail": "user2@test.com"}`,
- uid: 1,
- expectedCode: http.StatusConflict,
+ requestBody: `{"current_password": "SecurePass123!", "mail": "user2@test.com"}`,
+ uid: 1,
+ expectedCode: http.StatusConflict,
errorContains: "email already exists",
},
{
name: "invalid email format",
- requestBody: `{"current_password": "SecurePass123!", "mail": "invalid-email"}`,
- uid: 1,
- expectedCode: http.StatusBadRequest,
+ requestBody: `{"current_password": "SecurePass123!", "mail": "invalid-email"}`,
+ uid: 1,
+ expectedCode: http.StatusBadRequest,
errorContains: "failed to validate user email",
},
}
@@ -442,3 +445,10 @@ func TestChangeEmailCurrentUser(t *testing.T) {
}
}
+func TestIsUniqueViolation(t *testing.T) {
+ assert.True(t, isUniqueViolation(errors.New(`pq: duplicate key value violates unique constraint "users_mail_unique"`)))
+ assert.True(t, isUniqueViolation(errors.New("pq: error 23505")))
+ assert.True(t, isUniqueViolation(errors.New("UNIQUE constraint failed: users.mail")), "sqlite phrasing is matched case-insensitively")
+ assert.False(t, isUniqueViolation(errors.New("connection refused")))
+ assert.False(t, isUniqueViolation(nil))
+}
diff --git a/frontend/src/features/authentication/email-change-form.test.tsx b/frontend/src/features/authentication/email-change-form.test.tsx
index 12b8f640..0213b3f6 100644
--- a/frontend/src/features/authentication/email-change-form.test.tsx
+++ b/frontend/src/features/authentication/email-change-form.test.tsx
@@ -15,7 +15,7 @@ vi.mock('@/lib/axios', async (importOriginal) => {
});
vi.mock('@/providers/user-provider', () => ({
- useUser: () => ({ authInfo: { user: { mail: 'old@example.com' } }, refreshAuthInfo }),
+ useUser: () => ({ refreshAuthInfo }),
}));
vi.mock('sonner', () => ({ toast: { error: vi.fn(), success: vi.fn() } }));
@@ -30,21 +30,13 @@ beforeEach(() => {
});
describe('EmailChangeForm', () => {
- it('shows the current email in a label-associated disabled field', () => {
- render();
-
- const current = screen.getByLabelText('Current Email') as HTMLInputElement;
- expect(current).toBeDisabled();
- expect(current.value).toBe('old@example.com');
- });
-
it('submits the new email and refreshes auth before closing', async () => {
const user = userEvent.setup();
const onSuccess = vi.fn();
render();
- await user.type(screen.getByPlaceholderText('Enter new email address'), 'new@example.com');
- await user.type(screen.getByPlaceholderText('Enter your current password to confirm'), 'Oldpass0!');
+ await user.type(screen.getByPlaceholderText('Enter your new email address'), 'new@example.com');
+ await user.type(screen.getByPlaceholderText('Enter your current password'), 'Oldpass0!');
await user.click(screen.getByRole('button', { name: 'Update Email' }));
await waitFor(() => expect(onSuccess).toHaveBeenCalledOnce());
@@ -57,8 +49,8 @@ describe('EmailChangeForm', () => {
put.mockRejectedValueOnce(apiError('Users.ChangeEmailCurrentUser.EmailAlreadyExists', 'email already exists'));
render();
- await user.type(screen.getByPlaceholderText('Enter new email address'), 'taken@example.com');
- await user.type(screen.getByPlaceholderText('Enter your current password to confirm'), 'Oldpass0!');
+ await user.type(screen.getByPlaceholderText('Enter your new email address'), 'taken@example.com');
+ await user.type(screen.getByPlaceholderText('Enter your current password'), 'Oldpass0!');
await user.click(screen.getByRole('button', { name: 'Update Email' }));
expect(await screen.findByText('Email address is already in use')).toBeInTheDocument();
diff --git a/frontend/src/features/authentication/password-change-form.test.tsx b/frontend/src/features/authentication/password-change-form.test.tsx
index ed83b3d7..a4cce87e 100644
--- a/frontend/src/features/authentication/password-change-form.test.tsx
+++ b/frontend/src/features/authentication/password-change-form.test.tsx
@@ -40,8 +40,8 @@ describe('PasswordChangeForm', () => {
render();
await user.type(screen.getByPlaceholderText('Enter your current password'), 'Oldpass0!');
- await user.type(screen.getByPlaceholderText('Enter new password'), 'Abcdef1!gh');
- await user.type(screen.getByPlaceholderText('Confirm new password'), 'Abcdef1!gh');
+ await user.type(screen.getByPlaceholderText('Enter your new password'), 'Abcdef1!gh');
+ await user.type(screen.getByPlaceholderText('Confirm your new password'), 'Abcdef1!gh');
await user.click(screen.getByRole('button', { name: 'Update Password' }));
await waitFor(() => expect(onSuccess).toHaveBeenCalledOnce());
@@ -60,8 +60,8 @@ describe('PasswordChangeForm', () => {
render();
await user.type(screen.getByPlaceholderText('Enter your current password'), 'Oldpass0!');
- await user.type(screen.getByPlaceholderText('Enter new password'), 'Abcdef1!gh');
- await user.type(screen.getByPlaceholderText('Confirm new password'), 'Abcdef1!gh');
+ await user.type(screen.getByPlaceholderText('Enter your new password'), 'Abcdef1!gh');
+ await user.type(screen.getByPlaceholderText('Confirm your new password'), 'Abcdef1!gh');
await user.click(screen.getByRole('button', { name: 'Update Password' }));
expect(await screen.findByText('Current password is incorrect')).toBeInTheDocument();