Changeset 0.24.0 (#156)
This commit is contained in:
53
server/auth/auth.go
Normal file
53
server/auth/auth.go
Normal file
@@ -0,0 +1,53 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"errors"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/models"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/store"
|
||||
)
|
||||
|
||||
type Mode string
|
||||
|
||||
const (
|
||||
ModeBuiltin Mode = "builtin"
|
||||
ModeMTLS Mode = "mtls"
|
||||
ModeOIDC Mode = "oidc"
|
||||
)
|
||||
|
||||
func ParseMode(s string) (Mode, error) {
|
||||
switch Mode(s) {
|
||||
case ModeBuiltin, "":
|
||||
return ModeBuiltin, nil
|
||||
case ModeMTLS:
|
||||
return ModeMTLS, nil
|
||||
case ModeOIDC:
|
||||
return ModeOIDC, nil
|
||||
default:
|
||||
return "", ErrUnsupportedMode
|
||||
}
|
||||
}
|
||||
|
||||
type Result struct {
|
||||
User *models.User
|
||||
IsNewUser bool
|
||||
VaultHint string // password (builtin) or "" (external)
|
||||
}
|
||||
|
||||
type Provider interface {
|
||||
Mode() Mode
|
||||
Authenticate(c *gin.Context, stores store.Stores) (*Result, error)
|
||||
SupportsRegistration() bool
|
||||
Register(c *gin.Context, stores store.Stores) (*Result, error)
|
||||
}
|
||||
|
||||
var (
|
||||
ErrUnsupportedMode = errors.New("unsupported auth mode")
|
||||
ErrNotSupported = errors.New("operation not supported by this auth provider")
|
||||
ErrInvalidCreds = errors.New("invalid credentials")
|
||||
ErrInactive = errors.New("account is inactive")
|
||||
ErrRegistrationOff = errors.New("registration is disabled")
|
||||
ErrDuplicate = errors.New("username or email already taken")
|
||||
)
|
||||
130
server/auth/builtin.go
Normal file
130
server/auth/builtin.go
Normal file
@@ -0,0 +1,130 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/models"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/store"
|
||||
)
|
||||
|
||||
const bcryptCost = 12
|
||||
|
||||
type BuiltinProvider struct{}
|
||||
|
||||
func NewBuiltinProvider() *BuiltinProvider { return &BuiltinProvider{} }
|
||||
|
||||
func (p *BuiltinProvider) Mode() Mode { return ModeBuiltin }
|
||||
func (p *BuiltinProvider) SupportsRegistration() bool { return true }
|
||||
|
||||
func (p *BuiltinProvider) Authenticate(c *gin.Context, stores store.Stores) (*Result, error) {
|
||||
var req struct {
|
||||
Login string `json:"login" binding:"required"`
|
||||
Password string `json:"password" binding:"required"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
user, err := stores.Users.GetByLogin(c.Request.Context(), req.Login)
|
||||
if err != nil {
|
||||
return nil, ErrInvalidCreds
|
||||
}
|
||||
if !user.IsActive {
|
||||
return nil, ErrInactive
|
||||
}
|
||||
if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(req.Password)); err != nil {
|
||||
return nil, ErrInvalidCreds
|
||||
}
|
||||
|
||||
return &Result{User: user, VaultHint: req.Password}, nil
|
||||
}
|
||||
|
||||
func (p *BuiltinProvider) Register(c *gin.Context, stores store.Stores) (*Result, error) {
|
||||
var req struct {
|
||||
Username string `json:"username" binding:"required"`
|
||||
Email string `json:"email" binding:"required"`
|
||||
Password string `json:"password" binding:"required,min=8"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
ctx := c.Request.Context()
|
||||
|
||||
allowed, _ := stores.Policies.GetBool(ctx, "allow_registration")
|
||||
if !allowed {
|
||||
return nil, ErrRegistrationOff
|
||||
}
|
||||
|
||||
if existing, _ := stores.Users.GetByUsername(ctx, req.Username); existing != nil {
|
||||
return nil, ErrDuplicate
|
||||
}
|
||||
if existing, _ := stores.Users.GetByEmail(ctx, req.Email); existing != nil {
|
||||
return nil, ErrDuplicate
|
||||
}
|
||||
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(req.Password), bcryptCost)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
handle := UniqueHandle(ctx, stores.Users, models.HandleFromName(req.Username))
|
||||
|
||||
defaultActive, _ := stores.Policies.GetBool(ctx, "default_user_active")
|
||||
|
||||
user := &models.User{
|
||||
Username: strings.ToLower(req.Username),
|
||||
Email: strings.ToLower(req.Email),
|
||||
PasswordHash: string(hash),
|
||||
Role: models.UserRoleUser,
|
||||
IsActive: defaultActive,
|
||||
AuthSource: string(ModeBuiltin),
|
||||
Handle: handle,
|
||||
}
|
||||
|
||||
if err := stores.Users.Create(ctx, user); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &Result{User: user, IsNewUser: true, VaultHint: req.Password}, nil
|
||||
}
|
||||
|
||||
// UniqueHandle appends -2, -3, etc. on collision.
|
||||
func UniqueHandle(ctx context.Context, users store.UserStore, base string) string {
|
||||
if base == "" {
|
||||
base = "user"
|
||||
}
|
||||
if u, _ := users.GetByHandle(ctx, base); u == nil {
|
||||
return base
|
||||
}
|
||||
for i := 2; i < 1000; i++ {
|
||||
candidate := fmt.Sprintf("%s-%d", base, i)
|
||||
if u, _ := users.GetByHandle(ctx, candidate); u == nil {
|
||||
return candidate
|
||||
}
|
||||
}
|
||||
return fmt.Sprintf("%s-%s", base, store.NewID()[:8])
|
||||
}
|
||||
|
||||
func ErrorToHTTPStatus(err error) int {
|
||||
switch err {
|
||||
case ErrInvalidCreds:
|
||||
return http.StatusUnauthorized
|
||||
case ErrInactive:
|
||||
return http.StatusForbidden
|
||||
case ErrRegistrationOff:
|
||||
return http.StatusForbidden
|
||||
case ErrDuplicate:
|
||||
return http.StatusConflict
|
||||
case ErrNotSupported:
|
||||
return http.StatusMethodNotAllowed
|
||||
default:
|
||||
return http.StatusBadRequest
|
||||
}
|
||||
}
|
||||
@@ -67,6 +67,9 @@ type Config struct {
|
||||
// PROVIDER_AUTO_DISABLE_THRESHOLD: consecutive "down" hourly windows before a
|
||||
// provider is automatically deactivated. Default 3. Set to 0 to disable.
|
||||
ProviderAutoDisableThreshold int
|
||||
|
||||
// Auth mode (v0.24.0): "builtin" (default) | "mtls" | "oidc"
|
||||
AuthMode string
|
||||
}
|
||||
|
||||
// Load reads configuration from environment variables.
|
||||
@@ -105,6 +108,8 @@ func Load() *Config {
|
||||
WorkspaceIndexConcurrency: getEnvInt("WORKSPACE_INDEX_CONCURRENCY", 2),
|
||||
|
||||
ProviderAutoDisableThreshold: getEnvInt("PROVIDER_AUTO_DISABLE_THRESHOLD", 3),
|
||||
|
||||
AuthMode: getEnv("AUTH_MODE", "builtin"),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
13
server/database/migrations/018_auth_abstraction.sql
Normal file
13
server/database/migrations/018_auth_abstraction.sql
Normal file
@@ -0,0 +1,13 @@
|
||||
-- Chat Switchboard — 018 Auth Abstraction (v0.24.0)
|
||||
ALTER TABLE users
|
||||
ADD COLUMN IF NOT EXISTS auth_source VARCHAR(20) NOT NULL DEFAULT 'builtin'
|
||||
CHECK (auth_source IN ('builtin', 'mtls', 'oidc'));
|
||||
ALTER TABLE users ADD COLUMN IF NOT EXISTS external_id TEXT;
|
||||
ALTER TABLE users ADD COLUMN IF NOT EXISTS handle VARCHAR(100);
|
||||
ALTER TABLE users ALTER COLUMN password_hash DROP NOT NULL;
|
||||
UPDATE users SET handle = LOWER(REPLACE(REPLACE(username, ' ', '-'), '_', '-'))
|
||||
WHERE handle IS NULL;
|
||||
ALTER TABLE users ALTER COLUMN handle SET NOT NULL;
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_users_handle ON users(LOWER(handle));
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_users_external_id
|
||||
ON users(auth_source, external_id) WHERE external_id IS NOT NULL;
|
||||
@@ -0,0 +1,8 @@
|
||||
-- Chat Switchboard — 018 Auth Abstraction (SQLite) (v0.24.0)
|
||||
ALTER TABLE users ADD COLUMN auth_source TEXT NOT NULL DEFAULT 'builtin';
|
||||
ALTER TABLE users ADD COLUMN external_id TEXT;
|
||||
ALTER TABLE users ADD COLUMN handle TEXT;
|
||||
UPDATE users SET handle = LOWER(REPLACE(REPLACE(username, ' ', '-'), '_', '-'))
|
||||
WHERE handle IS NULL;
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_users_handle ON users(handle COLLATE NOCASE);
|
||||
CREATE UNIQUE INDEX IF NOT EXISTS idx_users_external_id ON users(auth_source, external_id);
|
||||
@@ -356,12 +356,13 @@ func TruncateAll(t *testing.T) {
|
||||
// SeedTestUser creates a test user and returns the user ID.
|
||||
func SeedTestUser(t *testing.T, username, email string) string {
|
||||
t.Helper()
|
||||
handle := strings.ToLower(strings.ReplaceAll(username, " ", "-"))
|
||||
if IsSQLite() {
|
||||
id := uuid.New().String()
|
||||
_, err := DB.Exec(`
|
||||
INSERT INTO users (id, username, email, password_hash, role)
|
||||
VALUES (?, ?, ?, '$2a$10$dummy.hash.for.testing.only.000000000000000000000', 'user')
|
||||
`, id, username, email)
|
||||
INSERT INTO users (id, username, email, password_hash, role, auth_source, handle)
|
||||
VALUES (?, ?, ?, '$2a$10$dummy.hash.for.testing.only.000000000000000000000', 'user', 'builtin', ?)
|
||||
`, id, username, email, handle)
|
||||
if err != nil {
|
||||
t.Fatalf("SeedTestUser: %v", err)
|
||||
}
|
||||
@@ -369,10 +370,10 @@ func SeedTestUser(t *testing.T, username, email string) string {
|
||||
}
|
||||
var id string
|
||||
err := DB.QueryRow(`
|
||||
INSERT INTO users (username, email, password_hash, role)
|
||||
VALUES ($1, $2, '$2a$10$dummy.hash.for.testing.only.000000000000000000000', 'user')
|
||||
INSERT INTO users (username, email, password_hash, role, auth_source, handle)
|
||||
VALUES ($1, $2, '$2a$10$dummy.hash.for.testing.only.000000000000000000000', 'user', 'builtin', $3)
|
||||
RETURNING id
|
||||
`, username, email).Scan(&id)
|
||||
`, username, email, handle).Scan(&id)
|
||||
if err != nil {
|
||||
t.Fatalf("SeedTestUser: %v", err)
|
||||
}
|
||||
|
||||
@@ -14,6 +14,7 @@ import (
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/auth"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/config"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/crypto"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/database"
|
||||
@@ -37,77 +38,40 @@ type AuthHandler struct {
|
||||
cfg *config.Config
|
||||
stores store.Stores
|
||||
uekCache *crypto.UEKCache
|
||||
provider auth.Provider
|
||||
}
|
||||
|
||||
func NewAuthHandler(cfg *config.Config, s store.Stores, uekCache *crypto.UEKCache) *AuthHandler {
|
||||
return &AuthHandler{cfg: cfg, stores: s, uekCache: uekCache}
|
||||
func NewAuthHandler(cfg *config.Config, s store.Stores, uekCache *crypto.UEKCache, provider auth.Provider) *AuthHandler {
|
||||
return &AuthHandler{cfg: cfg, stores: s, uekCache: uekCache, provider: provider}
|
||||
}
|
||||
|
||||
func (h *AuthHandler) Register(c *gin.Context) {
|
||||
var req struct {
|
||||
Username string `json:"username" binding:"required"`
|
||||
Email string `json:"email" binding:"required"`
|
||||
Password string `json:"password" binding:"required,min=8"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
if !h.provider.SupportsRegistration() {
|
||||
c.JSON(http.StatusMethodNotAllowed, gin.H{"error": "registration not supported for this auth mode"})
|
||||
return
|
||||
}
|
||||
|
||||
// Check registration policy
|
||||
allowed, _ := h.stores.Policies.GetBool(c.Request.Context(), "allow_registration")
|
||||
if !allowed {
|
||||
c.JSON(http.StatusForbidden, gin.H{"error": "registration is disabled"})
|
||||
return
|
||||
}
|
||||
|
||||
// Check duplicate
|
||||
if existing, _ := h.stores.Users.GetByUsername(c.Request.Context(), req.Username); existing != nil {
|
||||
c.JSON(http.StatusConflict, gin.H{"error": "username already taken"})
|
||||
return
|
||||
}
|
||||
if existing, _ := h.stores.Users.GetByEmail(c.Request.Context(), req.Email); existing != nil {
|
||||
c.JSON(http.StatusConflict, gin.H{"error": "email already registered"})
|
||||
return
|
||||
}
|
||||
|
||||
hash, err := bcrypt.GenerateFromPassword([]byte(req.Password), bcryptCost)
|
||||
result, err := h.provider.Register(c, h.stores)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to hash password"})
|
||||
c.JSON(auth.ErrorToHTTPStatus(err), gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
// Check if user should be active by default
|
||||
defaultActive, _ := h.stores.Policies.GetBool(c.Request.Context(), "default_user_active")
|
||||
|
||||
user := &models.User{
|
||||
Username: strings.ToLower(req.Username),
|
||||
Email: strings.ToLower(req.Email),
|
||||
PasswordHash: string(hash),
|
||||
Role: models.UserRoleUser,
|
||||
IsActive: defaultActive,
|
||||
if result.VaultHint != "" {
|
||||
if err := h.initVault(c.Request.Context(), result.User.ID, result.VaultHint); err != nil {
|
||||
log.Printf("⚠ Failed to init vault for user %s: %v", result.User.ID, err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := h.stores.Users.Create(c.Request.Context(), user); err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to create user"})
|
||||
return
|
||||
}
|
||||
|
||||
// Generate and store the User Encryption Key (vault)
|
||||
if err := h.initVault(c.Request.Context(), user.ID, req.Password); err != nil {
|
||||
log.Printf("⚠ Failed to init vault for user %s: %v", user.ID, err)
|
||||
// Non-fatal: user can still use the app, vault will init on next login
|
||||
}
|
||||
|
||||
if !user.IsActive {
|
||||
if !result.User.IsActive {
|
||||
c.JSON(http.StatusCreated, gin.H{
|
||||
"message": "Account created but requires admin approval",
|
||||
"user_id": user.ID,
|
||||
"user_id": result.User.ID,
|
||||
})
|
||||
return
|
||||
}
|
||||
|
||||
tokens, err := h.generateTokens(user)
|
||||
tokens, err := h.generateTokens(result.User)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to generate tokens"})
|
||||
return
|
||||
@@ -117,37 +81,19 @@ func (h *AuthHandler) Register(c *gin.Context) {
|
||||
}
|
||||
|
||||
func (h *AuthHandler) Login(c *gin.Context) {
|
||||
var req struct {
|
||||
Login string `json:"login" binding:"required"` // username or email
|
||||
Password string `json:"password" binding:"required"`
|
||||
}
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
user, err := h.stores.Users.GetByLogin(c.Request.Context(), req.Login)
|
||||
result, err := h.provider.Authenticate(c, h.stores)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"error": "invalid credentials"})
|
||||
c.JSON(auth.ErrorToHTTPStatus(err), gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
|
||||
if !user.IsActive {
|
||||
c.JSON(http.StatusForbidden, gin.H{"error": "account is inactive"})
|
||||
return
|
||||
if result.VaultHint != "" {
|
||||
h.unlockVault(c.Request.Context(), result.User, result.VaultHint)
|
||||
}
|
||||
|
||||
if err := bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(req.Password)); err != nil {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"error": "invalid credentials"})
|
||||
return
|
||||
}
|
||||
h.stores.Users.UpdateLastLogin(c.Request.Context(), result.User.ID)
|
||||
|
||||
// Vault: unwrap UEK into session cache (or init if first login post-migration)
|
||||
h.unlockVault(c.Request.Context(), user, req.Password)
|
||||
|
||||
h.stores.Users.UpdateLastLogin(c.Request.Context(), user.ID)
|
||||
|
||||
tokens, err := h.generateTokens(user)
|
||||
tokens, err := h.generateTokens(result.User)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to generate tokens"})
|
||||
return
|
||||
@@ -247,6 +193,8 @@ func (h *AuthHandler) generateTokens(user *models.User) (gin.H, error) {
|
||||
"email": user.Email,
|
||||
"display_name": user.DisplayName,
|
||||
"role": user.Role,
|
||||
"handle": user.Handle,
|
||||
"auth_source": user.AuthSource,
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
@@ -433,6 +381,11 @@ func BootstrapAdmin(cfg *config.Config, s store.Stores) {
|
||||
})
|
||||
// If the actual password changed, the vault seal is stale. Probe it
|
||||
// with the current password and destroy only if unwrap fails.
|
||||
// Ensure handle exists (backfill for pre-v0.24.0 users)
|
||||
if existing.Handle == "" {
|
||||
handle := auth.UniqueHandle(ctx, s.Users, models.HandleFromName(cfg.AdminUsername))
|
||||
s.Users.Update(ctx, existing.ID, map[string]interface{}{"handle": handle})
|
||||
}
|
||||
ProbeAndRepairVault(ctx, existing.ID, cfg.AdminPassword)
|
||||
log.Printf(" ✅ Admin user '%s' updated", cfg.AdminUsername)
|
||||
return
|
||||
@@ -443,12 +396,15 @@ func BootstrapAdmin(cfg *config.Config, s store.Stores) {
|
||||
email = cfg.AdminUsername + "@switchboard.local"
|
||||
}
|
||||
|
||||
handle := auth.UniqueHandle(ctx, s.Users, models.HandleFromName(cfg.AdminUsername))
|
||||
user := &models.User{
|
||||
Username: strings.ToLower(cfg.AdminUsername),
|
||||
Email: strings.ToLower(email),
|
||||
PasswordHash: string(hash),
|
||||
Role: models.UserRoleAdmin,
|
||||
IsActive: true,
|
||||
AuthSource: "builtin",
|
||||
Handle: handle,
|
||||
}
|
||||
|
||||
if err := s.Users.Create(ctx, user); err != nil {
|
||||
@@ -516,17 +472,24 @@ func SeedUsers(cfg *config.Config, s store.Stores) {
|
||||
"is_active": true,
|
||||
})
|
||||
// Probe vault with current password — only destroys if seal is stale
|
||||
if existing.Handle == "" {
|
||||
handle := auth.UniqueHandle(ctx, s.Users, models.HandleFromName(username))
|
||||
s.Users.Update(ctx, existing.ID, map[string]interface{}{"handle": handle})
|
||||
}
|
||||
ProbeAndRepairVault(ctx, existing.ID, password)
|
||||
log.Printf(" 🌱 Seed user '%s' updated (role=%s)", username, role)
|
||||
continue
|
||||
}
|
||||
|
||||
handle := auth.UniqueHandle(ctx, s.Users, models.HandleFromName(username))
|
||||
user := &models.User{
|
||||
Username: username,
|
||||
Email: username + "@switchboard.local",
|
||||
PasswordHash: string(hash),
|
||||
Role: role,
|
||||
IsActive: true,
|
||||
AuthSource: "builtin",
|
||||
Handle: handle,
|
||||
}
|
||||
|
||||
if err := s.Users.Create(ctx, user); err != nil {
|
||||
|
||||
@@ -14,6 +14,7 @@ import (
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/auth"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/config"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/store"
|
||||
)
|
||||
@@ -21,12 +22,13 @@ import (
|
||||
func testConfig() *config.Config {
|
||||
return &config.Config{
|
||||
JWTSecret: "test-secret-key-for-unit-tests",
|
||||
AuthMode: "builtin",
|
||||
}
|
||||
}
|
||||
|
||||
// testAuthHandler creates an AuthHandler with nil stores (safe for non-DB tests).
|
||||
func testAuthHandler() *AuthHandler {
|
||||
return NewAuthHandler(testConfig(), store.Stores{}, nil)
|
||||
return NewAuthHandler(testConfig(), store.Stores{}, nil, auth.NewBuiltinProvider())
|
||||
}
|
||||
|
||||
// ── JWT Token Tests ─────────────────────────
|
||||
|
||||
@@ -1340,12 +1340,12 @@ func (h *CompletionHandler) resolveMention(ctx context.Context, userID, content
|
||||
}
|
||||
}
|
||||
|
||||
// 3. Try username (exact match) — v0.23.1
|
||||
// 3. Try user handle (exact match) — v0.24.0
|
||||
// Only resolves to a user if the caller is not the same user (no self-mention)
|
||||
var mentionedUserID string
|
||||
err = database.DB.QueryRowContext(ctx, database.Q(`
|
||||
SELECT id FROM users
|
||||
WHERE LOWER(username) = $1 AND id != $2 AND is_active = true
|
||||
WHERE LOWER(handle) = $1 AND id != $2 AND is_active = true
|
||||
LIMIT 1
|
||||
`), normalized, userID).Scan(&mentionedUserID)
|
||||
if err == nil && mentionedUserID != "" {
|
||||
@@ -1353,16 +1353,16 @@ func (h *CompletionHandler) resolveMention(ctx context.Context, userID, content
|
||||
return "", "", nil, mentionedUserID
|
||||
}
|
||||
|
||||
// 4. Try username (prefix — unambiguous only) — v0.23.1
|
||||
// 4. Try user handle (prefix — unambiguous only) — v0.24.0
|
||||
var userCount int
|
||||
database.DB.QueryRowContext(ctx, database.Q(`
|
||||
SELECT COUNT(*) FROM users
|
||||
WHERE LOWER(username) LIKE $1 AND id != $2 AND is_active = true
|
||||
WHERE LOWER(handle) LIKE $1 AND id != $2 AND is_active = true
|
||||
`), normalized+"%", userID).Scan(&userCount)
|
||||
if userCount == 1 {
|
||||
database.DB.QueryRowContext(ctx, database.Q(`
|
||||
SELECT id FROM users
|
||||
WHERE LOWER(username) LIKE $1 AND id != $2 AND is_active = true
|
||||
WHERE LOWER(handle) LIKE $1 AND id != $2 AND is_active = true
|
||||
LIMIT 1
|
||||
`), normalized+"%", userID).Scan(&mentionedUserID)
|
||||
if mentionedUserID != "" {
|
||||
|
||||
@@ -16,6 +16,7 @@ import (
|
||||
"github.com/google/uuid"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/config"
|
||||
authpkg "git.gobha.me/xcaliber/chat-switchboard/auth"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/database"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/middleware"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/models"
|
||||
@@ -135,7 +136,7 @@ func setupHarness(t *testing.T) *testHarness {
|
||||
api := r.Group("/api/v1")
|
||||
|
||||
// Auth (unprotected)
|
||||
auth := NewAuthHandler(cfg, stores, nil)
|
||||
auth := NewAuthHandler(cfg, stores, nil, authpkg.NewBuiltinProvider())
|
||||
authGroup := api.Group("/auth")
|
||||
authGroup.POST("/register", auth.Register)
|
||||
authGroup.POST("/login", auth.Login)
|
||||
|
||||
@@ -76,16 +76,16 @@ func SearchUsers(c *gin.Context) {
|
||||
q := strings.TrimSpace(c.Query("q"))
|
||||
|
||||
query := database.Q(`
|
||||
SELECT id, username, COALESCE(display_name, '') AS display_name
|
||||
SELECT id, username, COALESCE(display_name, '') AS display_name, COALESCE(handle, '') AS handle
|
||||
FROM users
|
||||
WHERE is_active = true AND id != $1
|
||||
`)
|
||||
args := []interface{}{userID}
|
||||
|
||||
if q != "" {
|
||||
query += database.Q(` AND (LOWER(username) LIKE $2 OR LOWER(display_name) LIKE $3)`)
|
||||
query += database.Q(` AND (LOWER(username) LIKE $2 OR LOWER(display_name) LIKE $3 OR LOWER(handle) LIKE $4)`)
|
||||
pattern := "%" + strings.ToLower(q) + "%"
|
||||
args = append(args, pattern, pattern)
|
||||
args = append(args, pattern, pattern, pattern)
|
||||
}
|
||||
|
||||
query += ` ORDER BY username LIMIT 20`
|
||||
@@ -101,12 +101,13 @@ func SearchUsers(c *gin.Context) {
|
||||
ID string `json:"id"`
|
||||
Username string `json:"username"`
|
||||
DisplayName string `json:"display_name"`
|
||||
Handle string `json:"handle"`
|
||||
}
|
||||
|
||||
results := []userResult{}
|
||||
for rows.Next() {
|
||||
var u userResult
|
||||
if err := rows.Scan(&u.ID, &u.Username, &u.DisplayName); err != nil {
|
||||
if err := rows.Scan(&u.ID, &u.Username, &u.DisplayName, &u.Handle); err != nil {
|
||||
continue
|
||||
}
|
||||
results = append(results, u)
|
||||
|
||||
@@ -11,6 +11,7 @@ import (
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"git.gobha.me/xcaliber/chat-switchboard/auth"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/compaction"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/config"
|
||||
"git.gobha.me/xcaliber/chat-switchboard/crypto"
|
||||
@@ -329,7 +330,22 @@ func main() {
|
||||
base.GET("/ws", middleware.Auth(cfg), hub.HandleWebSocket)
|
||||
|
||||
// ── Auth routes (rate limited) ──────────────
|
||||
auth := handlers.NewAuthHandler(cfg, stores, uekCache)
|
||||
authMode, err := auth.ParseMode(cfg.AuthMode)
|
||||
if err != nil {
|
||||
log.Fatalf("❌ Invalid AUTH_MODE=%q: %v", cfg.AuthMode, err)
|
||||
}
|
||||
var authProvider auth.Provider
|
||||
switch authMode {
|
||||
case auth.ModeBuiltin:
|
||||
authProvider = auth.NewBuiltinProvider()
|
||||
case auth.ModeMTLS:
|
||||
log.Fatal("❌ AUTH_MODE=mtls is not yet implemented (planned for v0.24.1)")
|
||||
case auth.ModeOIDC:
|
||||
log.Fatal("❌ AUTH_MODE=oidc is not yet implemented (planned for v0.24.1)")
|
||||
}
|
||||
log.Printf(" 🔑 Auth mode: %s", authMode)
|
||||
|
||||
authH := handlers.NewAuthHandler(cfg, stores, uekCache, authProvider)
|
||||
authLimiter := middleware.NewRateLimiter(1, 5)
|
||||
|
||||
api := base.Group("/api/v1")
|
||||
@@ -353,10 +369,10 @@ func main() {
|
||||
authGroup := api.Group("/auth")
|
||||
authGroup.Use(authLimiter.Limit())
|
||||
{
|
||||
authGroup.POST("/register", auth.Register)
|
||||
authGroup.POST("/login", auth.Login)
|
||||
authGroup.POST("/refresh", auth.Refresh)
|
||||
authGroup.POST("/logout", auth.Logout)
|
||||
authGroup.POST("/register", authH.Register)
|
||||
authGroup.POST("/login", authH.Login)
|
||||
authGroup.POST("/refresh", authH.Refresh)
|
||||
authGroup.POST("/logout", authH.Logout)
|
||||
}
|
||||
|
||||
// ── Public extension assets ────────────────
|
||||
|
||||
@@ -80,6 +80,9 @@ type User struct {
|
||||
IsActive bool `json:"is_active" db:"is_active"`
|
||||
Settings JSONMap `json:"settings,omitempty" db:"settings"`
|
||||
LastLoginAt *time.Time `json:"last_login_at,omitempty" db:"last_login_at"`
|
||||
AuthSource string `json:"auth_source" db:"auth_source"`
|
||||
ExternalID *string `json:"external_id,omitempty" db:"external_id"`
|
||||
Handle string `json:"handle" db:"handle"`
|
||||
}
|
||||
|
||||
// =========================================
|
||||
|
||||
@@ -167,6 +167,8 @@ type UserStore interface {
|
||||
GetByUsername(ctx context.Context, username string) (*models.User, error)
|
||||
GetByEmail(ctx context.Context, email string) (*models.User, error)
|
||||
GetByLogin(ctx context.Context, login string) (*models.User, error) // username or email
|
||||
GetByHandle(ctx context.Context, handle string) (*models.User, error)
|
||||
GetByExternalID(ctx context.Context, authSource, externalID string) (*models.User, error)
|
||||
Update(ctx context.Context, id string, fields map[string]interface{}) error
|
||||
Delete(ctx context.Context, id string) error
|
||||
List(ctx context.Context, opts ListOptions) ([]models.User, int, error)
|
||||
|
||||
@@ -14,77 +14,43 @@ type UserStore struct{}
|
||||
|
||||
func NewUserStore() *UserStore { return &UserStore{} }
|
||||
|
||||
const userCols = `id, username, email, password_hash, display_name, avatar_url,
|
||||
role, is_active, settings, created_at, updated_at, last_login_at,
|
||||
auth_source, external_id, handle`
|
||||
|
||||
const userListCols = `id, username, email, display_name, avatar_url, role, is_active,
|
||||
settings, created_at, updated_at, last_login_at, auth_source, external_id, handle`
|
||||
|
||||
func (s *UserStore) Create(ctx context.Context, u *models.User) error {
|
||||
if u.AuthSource == "" {
|
||||
u.AuthSource = "builtin"
|
||||
}
|
||||
return DB.QueryRowContext(ctx, `
|
||||
INSERT INTO users (username, email, password_hash, display_name, role, is_active, settings)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7)
|
||||
INSERT INTO users (username, email, password_hash, display_name, role, is_active, settings, auth_source, external_id, handle)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10)
|
||||
RETURNING id, created_at, updated_at`,
|
||||
u.Username, u.Email, u.PasswordHash, u.DisplayName, u.Role, u.IsActive, ToJSON(u.Settings),
|
||||
u.Username, u.Email, u.PasswordHash, u.DisplayName, u.Role, u.IsActive,
|
||||
ToJSON(u.Settings), u.AuthSource, u.ExternalID, u.Handle,
|
||||
).Scan(&u.ID, &u.CreatedAt, &u.UpdatedAt)
|
||||
}
|
||||
|
||||
func (s *UserStore) GetByID(ctx context.Context, id string) (*models.User, error) {
|
||||
return s.getBy(ctx, "id", id)
|
||||
return scanOneUser(ctx, fmt.Sprintf("SELECT %s FROM users WHERE id = $1", userCols), id)
|
||||
}
|
||||
|
||||
func (s *UserStore) GetByUsername(ctx context.Context, username string) (*models.User, error) {
|
||||
var u models.User
|
||||
var displayName, avatarURL sql.NullString
|
||||
var settingsJSON []byte
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
SELECT id, username, email, password_hash, display_name, avatar_url,
|
||||
role, is_active, settings, created_at, updated_at, last_login_at
|
||||
FROM users WHERE LOWER(username) = LOWER($1)`, username).Scan(
|
||||
&u.ID, &u.Username, &u.Email, &u.PasswordHash, &displayName, &avatarURL,
|
||||
&u.Role, &u.IsActive, &settingsJSON, &u.CreatedAt, &u.UpdatedAt, &u.LastLoginAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
u.DisplayName = NullableString(displayName)
|
||||
u.AvatarURL = NullableString(avatarURL)
|
||||
ScanJSON(settingsJSON, &u.Settings)
|
||||
return &u, nil
|
||||
return scanOneUser(ctx, fmt.Sprintf("SELECT %s FROM users WHERE LOWER(username) = LOWER($1)", userCols), username)
|
||||
}
|
||||
|
||||
func (s *UserStore) GetByEmail(ctx context.Context, email string) (*models.User, error) {
|
||||
var u models.User
|
||||
var displayName, avatarURL sql.NullString
|
||||
var settingsJSON []byte
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
SELECT id, username, email, password_hash, display_name, avatar_url,
|
||||
role, is_active, settings, created_at, updated_at, last_login_at
|
||||
FROM users WHERE LOWER(email) = LOWER($1)`, email).Scan(
|
||||
&u.ID, &u.Username, &u.Email, &u.PasswordHash, &displayName, &avatarURL,
|
||||
&u.Role, &u.IsActive, &settingsJSON, &u.CreatedAt, &u.UpdatedAt, &u.LastLoginAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
u.DisplayName = NullableString(displayName)
|
||||
u.AvatarURL = NullableString(avatarURL)
|
||||
ScanJSON(settingsJSON, &u.Settings)
|
||||
return &u, nil
|
||||
return scanOneUser(ctx, fmt.Sprintf("SELECT %s FROM users WHERE LOWER(email) = LOWER($1)", userCols), email)
|
||||
}
|
||||
|
||||
func (s *UserStore) GetByLogin(ctx context.Context, login string) (*models.User, error) {
|
||||
var u models.User
|
||||
var displayName, avatarURL sql.NullString
|
||||
var settingsJSON []byte
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
SELECT id, username, email, password_hash, display_name, avatar_url,
|
||||
role, is_active, settings, created_at, updated_at, last_login_at
|
||||
FROM users WHERE LOWER(username) = LOWER($1) OR LOWER(email) = LOWER($1)`, login).Scan(
|
||||
&u.ID, &u.Username, &u.Email, &u.PasswordHash, &displayName, &avatarURL,
|
||||
&u.Role, &u.IsActive, &settingsJSON, &u.CreatedAt, &u.UpdatedAt, &u.LastLoginAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
u.DisplayName = NullableString(displayName)
|
||||
u.AvatarURL = NullableString(avatarURL)
|
||||
ScanJSON(settingsJSON, &u.Settings)
|
||||
return &u, nil
|
||||
return scanOneUser(ctx, fmt.Sprintf("SELECT %s FROM users WHERE LOWER(username) = LOWER($1) OR LOWER(email) = LOWER($1)", userCols), login)
|
||||
}
|
||||
func (s *UserStore) GetByHandle(ctx context.Context, handle string) (*models.User, error) {
|
||||
return scanOneUser(ctx, fmt.Sprintf("SELECT %s FROM users WHERE LOWER(handle) = LOWER($1)", userCols), handle)
|
||||
}
|
||||
func (s *UserStore) GetByExternalID(ctx context.Context, authSource, externalID string) (*models.User, error) {
|
||||
return scanOneUser(ctx, fmt.Sprintf("SELECT %s FROM users WHERE auth_source = $1 AND external_id = $2", userCols), authSource, externalID)
|
||||
}
|
||||
|
||||
func (s *UserStore) Update(ctx context.Context, id string, fields map[string]interface{}) error {
|
||||
@@ -110,16 +76,12 @@ func (s *UserStore) Delete(ctx context.Context, id string) error {
|
||||
}
|
||||
|
||||
func (s *UserStore) List(ctx context.Context, opts store.ListOptions) ([]models.User, int, error) {
|
||||
b := NewSelect(
|
||||
"id, username, email, display_name, avatar_url, role, is_active, settings, created_at, updated_at, last_login_at",
|
||||
"users",
|
||||
)
|
||||
b := NewSelect(userListCols, "users")
|
||||
if opts.Sort == "" {
|
||||
b.OrderBy("username", "ASC")
|
||||
}
|
||||
b.Paginate(opts)
|
||||
|
||||
// Count
|
||||
var total int
|
||||
DB.QueryRowContext(ctx, "SELECT COUNT(*) FROM users").Scan(&total)
|
||||
|
||||
@@ -133,16 +95,19 @@ func (s *UserStore) List(ctx context.Context, opts store.ListOptions) ([]models.
|
||||
var result []models.User
|
||||
for rows.Next() {
|
||||
var u models.User
|
||||
var displayName, avatarURL sql.NullString
|
||||
var settingsJSON []byte
|
||||
err := rows.Scan(&u.ID, &u.Username, &u.Email, &displayName, &avatarURL,
|
||||
&u.Role, &u.IsActive, &settingsJSON, &u.CreatedAt, &u.UpdatedAt, &u.LastLoginAt)
|
||||
if err != nil {
|
||||
var dn, av sql.NullString
|
||||
var extID, hdl sql.NullString
|
||||
var sj []byte
|
||||
if err := rows.Scan(&u.ID, &u.Username, &u.Email, &dn, &av,
|
||||
&u.Role, &u.IsActive, &sj, &u.CreatedAt, &u.UpdatedAt, &u.LastLoginAt,
|
||||
&u.AuthSource, &extID, &hdl); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
u.DisplayName = NullableString(displayName)
|
||||
u.AvatarURL = NullableString(avatarURL)
|
||||
ScanJSON(settingsJSON, &u.Settings)
|
||||
u.DisplayName = NullableString(dn)
|
||||
u.AvatarURL = NullableString(av)
|
||||
if extID.Valid { u.ExternalID = &extID.String }
|
||||
u.Handle = NullableString(hdl)
|
||||
ScanJSON(sj, &u.Settings)
|
||||
result = append(result, u)
|
||||
}
|
||||
return result, total, rows.Err()
|
||||
@@ -152,7 +117,6 @@ func (s *UserStore) UpdateLastLogin(ctx context.Context, id string) error {
|
||||
_, err := DB.ExecContext(ctx, "UPDATE users SET last_login_at = NOW() WHERE id = $1", id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *UserStore) SetActive(ctx context.Context, id string, active bool) error {
|
||||
_, err := DB.ExecContext(ctx, "UPDATE users SET is_active = $1 WHERE id = $2", active, id)
|
||||
return err
|
||||
@@ -161,57 +125,47 @@ func (s *UserStore) SetActive(ctx context.Context, id string, active bool) error
|
||||
// ── Refresh Tokens ──────────────────────────
|
||||
|
||||
func (s *UserStore) CreateRefreshToken(ctx context.Context, userID, tokenHash string, expiresAt time.Time) error {
|
||||
_, err := DB.ExecContext(ctx, `
|
||||
INSERT INTO refresh_tokens (user_id, token_hash, expires_at)
|
||||
VALUES ($1, $2, $3)`, userID, tokenHash, expiresAt)
|
||||
_, err := DB.ExecContext(ctx, `INSERT INTO refresh_tokens (user_id, token_hash, expires_at) VALUES ($1, $2, $3)`, userID, tokenHash, expiresAt)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *UserStore) GetRefreshToken(ctx context.Context, tokenHash string) (string, error) {
|
||||
var userID string
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
SELECT user_id FROM refresh_tokens
|
||||
WHERE token_hash = $1 AND revoked_at IS NULL AND expires_at > NOW()`,
|
||||
tokenHash).Scan(&userID)
|
||||
err := DB.QueryRowContext(ctx, `SELECT user_id FROM refresh_tokens WHERE token_hash = $1 AND revoked_at IS NULL AND expires_at > NOW()`, tokenHash).Scan(&userID)
|
||||
return userID, err
|
||||
}
|
||||
|
||||
func (s *UserStore) RevokeRefreshToken(ctx context.Context, tokenHash string) error {
|
||||
_, err := DB.ExecContext(ctx,
|
||||
"UPDATE refresh_tokens SET revoked_at = NOW() WHERE token_hash = $1", tokenHash)
|
||||
_, err := DB.ExecContext(ctx, "UPDATE refresh_tokens SET revoked_at = NOW() WHERE token_hash = $1", tokenHash)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *UserStore) RevokeAllRefreshTokens(ctx context.Context, userID string) error {
|
||||
_, err := DB.ExecContext(ctx,
|
||||
"UPDATE refresh_tokens SET revoked_at = NOW() WHERE user_id = $1 AND revoked_at IS NULL", userID)
|
||||
_, err := DB.ExecContext(ctx, "UPDATE refresh_tokens SET revoked_at = NOW() WHERE user_id = $1 AND revoked_at IS NULL", userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *UserStore) CleanExpiredTokens(ctx context.Context) error {
|
||||
_, err := DB.ExecContext(ctx,
|
||||
"DELETE FROM refresh_tokens WHERE expires_at < NOW() - INTERVAL '30 days'")
|
||||
_, err := DB.ExecContext(ctx, "DELETE FROM refresh_tokens WHERE expires_at < NOW() - INTERVAL '30 days'")
|
||||
return err
|
||||
}
|
||||
|
||||
// ── Internal ────────────────────────────────
|
||||
|
||||
func (s *UserStore) getBy(ctx context.Context, col, val string) (*models.User, error) {
|
||||
func scanOneUser(ctx context.Context, query string, args ...interface{}) (*models.User, error) {
|
||||
var u models.User
|
||||
var displayName, avatarURL sql.NullString
|
||||
var settingsJSON []byte
|
||||
err := DB.QueryRowContext(ctx, fmt.Sprintf(`
|
||||
SELECT id, username, email, password_hash, display_name, avatar_url,
|
||||
role, is_active, settings, created_at, updated_at, last_login_at
|
||||
FROM users WHERE %s = $1`, col), val).Scan(
|
||||
&u.ID, &u.Username, &u.Email, &u.PasswordHash, &displayName, &avatarURL,
|
||||
&u.Role, &u.IsActive, &settingsJSON, &u.CreatedAt, &u.UpdatedAt, &u.LastLoginAt,
|
||||
var dn, av, ph sql.NullString
|
||||
var extID, hdl sql.NullString
|
||||
var sj []byte
|
||||
err := DB.QueryRowContext(ctx, query, args...).Scan(
|
||||
&u.ID, &u.Username, &u.Email, &ph, &dn, &av,
|
||||
&u.Role, &u.IsActive, &sj, &u.CreatedAt, &u.UpdatedAt, &u.LastLoginAt,
|
||||
&u.AuthSource, &extID, &hdl,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
u.DisplayName = NullableString(displayName)
|
||||
u.AvatarURL = NullableString(avatarURL)
|
||||
ScanJSON(settingsJSON, &u.Settings)
|
||||
u.PasswordHash = NullableString(ph)
|
||||
u.DisplayName = NullableString(dn)
|
||||
u.AvatarURL = NullableString(av)
|
||||
if extID.Valid { u.ExternalID = &extID.String }
|
||||
u.Handle = NullableString(hdl)
|
||||
ScanJSON(sj, &u.Settings)
|
||||
return &u, nil
|
||||
}
|
||||
|
||||
@@ -14,82 +14,49 @@ type UserStore struct{}
|
||||
|
||||
func NewUserStore() *UserStore { return &UserStore{} }
|
||||
|
||||
const userCols = `id, username, email, password_hash, display_name, avatar_url,
|
||||
role, is_active, settings, created_at, updated_at, last_login_at,
|
||||
auth_source, external_id, handle`
|
||||
|
||||
const userListCols = `id, username, email, display_name, avatar_url, role, is_active,
|
||||
settings, created_at, updated_at, last_login_at, auth_source, external_id, handle`
|
||||
|
||||
func (s *UserStore) Create(ctx context.Context, u *models.User) error {
|
||||
u.ID = store.NewID()
|
||||
now := time.Now().UTC()
|
||||
u.CreatedAt = now
|
||||
u.UpdatedAt = now
|
||||
if u.AuthSource == "" {
|
||||
u.AuthSource = "builtin"
|
||||
}
|
||||
_, err := DB.ExecContext(ctx, `
|
||||
INSERT INTO users (id, username, email, password_hash, display_name, role, is_active, settings, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
|
||||
u.ID, u.Username, u.Email, u.PasswordHash, u.DisplayName, u.Role, u.IsActive, ToJSON(u.Settings),
|
||||
now.Format(timeFmt), now.Format(timeFmt),
|
||||
INSERT INTO users (id, username, email, password_hash, display_name, role, is_active,
|
||||
settings, created_at, updated_at, auth_source, external_id, handle)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
|
||||
u.ID, u.Username, u.Email, u.PasswordHash, u.DisplayName, u.Role, u.IsActive,
|
||||
ToJSON(u.Settings), now.Format(timeFmt), now.Format(timeFmt),
|
||||
u.AuthSource, u.ExternalID, u.Handle,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *UserStore) GetByID(ctx context.Context, id string) (*models.User, error) {
|
||||
return s.getBy(ctx, "id", id)
|
||||
return s.scanOne(ctx, fmt.Sprintf("SELECT %s FROM users WHERE id = ?", userCols), id)
|
||||
}
|
||||
|
||||
func (s *UserStore) GetByUsername(ctx context.Context, username string) (*models.User, error) {
|
||||
var u models.User
|
||||
var displayName, avatarURL sql.NullString
|
||||
var settingsJSON []byte
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
SELECT id, username, email, password_hash, display_name, avatar_url,
|
||||
role, is_active, settings, created_at, updated_at, last_login_at
|
||||
FROM users WHERE LOWER(username) = LOWER(?)`, username).Scan(
|
||||
&u.ID, &u.Username, &u.Email, &u.PasswordHash, &displayName, &avatarURL,
|
||||
&u.Role, &u.IsActive, &settingsJSON, st(&u.CreatedAt), st(&u.UpdatedAt), stN(&u.LastLoginAt),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
u.DisplayName = NullableString(displayName)
|
||||
u.AvatarURL = NullableString(avatarURL)
|
||||
ScanJSON(settingsJSON, &u.Settings)
|
||||
return &u, nil
|
||||
return s.scanOne(ctx, fmt.Sprintf("SELECT %s FROM users WHERE LOWER(username) = LOWER(?)", userCols), username)
|
||||
}
|
||||
|
||||
func (s *UserStore) GetByEmail(ctx context.Context, email string) (*models.User, error) {
|
||||
var u models.User
|
||||
var displayName, avatarURL sql.NullString
|
||||
var settingsJSON []byte
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
SELECT id, username, email, password_hash, display_name, avatar_url,
|
||||
role, is_active, settings, created_at, updated_at, last_login_at
|
||||
FROM users WHERE LOWER(email) = LOWER(?)`, email).Scan(
|
||||
&u.ID, &u.Username, &u.Email, &u.PasswordHash, &displayName, &avatarURL,
|
||||
&u.Role, &u.IsActive, &settingsJSON, st(&u.CreatedAt), st(&u.UpdatedAt), stN(&u.LastLoginAt),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
u.DisplayName = NullableString(displayName)
|
||||
u.AvatarURL = NullableString(avatarURL)
|
||||
ScanJSON(settingsJSON, &u.Settings)
|
||||
return &u, nil
|
||||
return s.scanOne(ctx, fmt.Sprintf("SELECT %s FROM users WHERE LOWER(email) = LOWER(?)", userCols), email)
|
||||
}
|
||||
|
||||
func (s *UserStore) GetByLogin(ctx context.Context, login string) (*models.User, error) {
|
||||
var u models.User
|
||||
var displayName, avatarURL sql.NullString
|
||||
var settingsJSON []byte
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
SELECT id, username, email, password_hash, display_name, avatar_url,
|
||||
role, is_active, settings, created_at, updated_at, last_login_at
|
||||
FROM users WHERE LOWER(username) = LOWER(?) OR LOWER(email) = LOWER(?)`, login, login).Scan(
|
||||
&u.ID, &u.Username, &u.Email, &u.PasswordHash, &displayName, &avatarURL,
|
||||
&u.Role, &u.IsActive, &settingsJSON, st(&u.CreatedAt), st(&u.UpdatedAt), stN(&u.LastLoginAt),
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
u.DisplayName = NullableString(displayName)
|
||||
u.AvatarURL = NullableString(avatarURL)
|
||||
ScanJSON(settingsJSON, &u.Settings)
|
||||
return &u, nil
|
||||
return s.scanOne(ctx, fmt.Sprintf("SELECT %s FROM users WHERE LOWER(username) = LOWER(?) OR LOWER(email) = LOWER(?)", userCols), login, login)
|
||||
}
|
||||
func (s *UserStore) GetByHandle(ctx context.Context, handle string) (*models.User, error) {
|
||||
return s.scanOne(ctx, fmt.Sprintf("SELECT %s FROM users WHERE LOWER(handle) = LOWER(?)", userCols), handle)
|
||||
}
|
||||
func (s *UserStore) GetByExternalID(ctx context.Context, authSource, externalID string) (*models.User, error) {
|
||||
return s.scanOne(ctx, fmt.Sprintf("SELECT %s FROM users WHERE auth_source = ? AND external_id = ?", userCols), authSource, externalID)
|
||||
}
|
||||
|
||||
func (s *UserStore) Update(ctx context.Context, id string, fields map[string]interface{}) error {
|
||||
@@ -115,16 +82,12 @@ func (s *UserStore) Delete(ctx context.Context, id string) error {
|
||||
}
|
||||
|
||||
func (s *UserStore) List(ctx context.Context, opts store.ListOptions) ([]models.User, int, error) {
|
||||
b := NewSelect(
|
||||
"id, username, email, display_name, avatar_url, role, is_active, settings, created_at, updated_at, last_login_at",
|
||||
"users",
|
||||
)
|
||||
b := NewSelect(userListCols, "users")
|
||||
if opts.Sort == "" {
|
||||
b.OrderBy("username", "ASC")
|
||||
}
|
||||
b.Paginate(opts)
|
||||
|
||||
// Count
|
||||
var total int
|
||||
DB.QueryRowContext(ctx, "SELECT COUNT(*) FROM users").Scan(&total)
|
||||
|
||||
@@ -138,16 +101,19 @@ func (s *UserStore) List(ctx context.Context, opts store.ListOptions) ([]models.
|
||||
var result []models.User
|
||||
for rows.Next() {
|
||||
var u models.User
|
||||
var displayName, avatarURL sql.NullString
|
||||
var settingsJSON []byte
|
||||
err := rows.Scan(&u.ID, &u.Username, &u.Email, &displayName, &avatarURL,
|
||||
&u.Role, &u.IsActive, &settingsJSON, st(&u.CreatedAt), st(&u.UpdatedAt), stN(&u.LastLoginAt))
|
||||
if err != nil {
|
||||
var dn, av sql.NullString
|
||||
var extID, hdl sql.NullString
|
||||
var sj []byte
|
||||
if err := rows.Scan(&u.ID, &u.Username, &u.Email, &dn, &av,
|
||||
&u.Role, &u.IsActive, &sj, st(&u.CreatedAt), st(&u.UpdatedAt), stN(&u.LastLoginAt),
|
||||
&u.AuthSource, &extID, &hdl); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
u.DisplayName = NullableString(displayName)
|
||||
u.AvatarURL = NullableString(avatarURL)
|
||||
ScanJSON(settingsJSON, &u.Settings)
|
||||
u.DisplayName = NullableString(dn)
|
||||
u.AvatarURL = NullableString(av)
|
||||
if extID.Valid { u.ExternalID = &extID.String }
|
||||
u.Handle = NullableString(hdl)
|
||||
ScanJSON(sj, &u.Settings)
|
||||
result = append(result, u)
|
||||
}
|
||||
return result, total, rows.Err()
|
||||
@@ -157,7 +123,6 @@ func (s *UserStore) UpdateLastLogin(ctx context.Context, id string) error {
|
||||
_, err := DB.ExecContext(ctx, "UPDATE users SET last_login_at = datetime('now') WHERE id = ?", id)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *UserStore) SetActive(ctx context.Context, id string, active bool) error {
|
||||
_, err := DB.ExecContext(ctx, "UPDATE users SET is_active = ? WHERE id = ?", active, id)
|
||||
return err
|
||||
@@ -166,57 +131,48 @@ func (s *UserStore) SetActive(ctx context.Context, id string, active bool) error
|
||||
// ── Refresh Tokens ──────────────────────────
|
||||
|
||||
func (s *UserStore) CreateRefreshToken(ctx context.Context, userID, tokenHash string, expiresAt time.Time) error {
|
||||
_, err := DB.ExecContext(ctx, `
|
||||
INSERT INTO refresh_tokens (id, user_id, token_hash, expires_at)
|
||||
VALUES (?, ?, ?, ?)`, store.NewID(), userID, tokenHash, expiresAt.Format(timeFmt))
|
||||
_, err := DB.ExecContext(ctx, `INSERT INTO refresh_tokens (id, user_id, token_hash, expires_at) VALUES (?, ?, ?, ?)`,
|
||||
store.NewID(), userID, tokenHash, expiresAt.Format(timeFmt))
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *UserStore) GetRefreshToken(ctx context.Context, tokenHash string) (string, error) {
|
||||
var userID string
|
||||
err := DB.QueryRowContext(ctx, `
|
||||
SELECT user_id FROM refresh_tokens
|
||||
WHERE token_hash = ? AND revoked_at IS NULL AND expires_at > datetime('now')`,
|
||||
tokenHash).Scan(&userID)
|
||||
err := DB.QueryRowContext(ctx, `SELECT user_id FROM refresh_tokens WHERE token_hash = ? AND revoked_at IS NULL AND expires_at > datetime('now')`, tokenHash).Scan(&userID)
|
||||
return userID, err
|
||||
}
|
||||
|
||||
func (s *UserStore) RevokeRefreshToken(ctx context.Context, tokenHash string) error {
|
||||
_, err := DB.ExecContext(ctx,
|
||||
"UPDATE refresh_tokens SET revoked_at = datetime('now') WHERE token_hash = ?", tokenHash)
|
||||
_, err := DB.ExecContext(ctx, "UPDATE refresh_tokens SET revoked_at = datetime('now') WHERE token_hash = ?", tokenHash)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *UserStore) RevokeAllRefreshTokens(ctx context.Context, userID string) error {
|
||||
_, err := DB.ExecContext(ctx,
|
||||
"UPDATE refresh_tokens SET revoked_at = datetime('now') WHERE user_id = ? AND revoked_at IS NULL", userID)
|
||||
_, err := DB.ExecContext(ctx, "UPDATE refresh_tokens SET revoked_at = datetime('now') WHERE user_id = ? AND revoked_at IS NULL", userID)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *UserStore) CleanExpiredTokens(ctx context.Context) error {
|
||||
_, err := DB.ExecContext(ctx,
|
||||
"DELETE FROM refresh_tokens WHERE expires_at < datetime('now', '-30 days')")
|
||||
_, err := DB.ExecContext(ctx, "DELETE FROM refresh_tokens WHERE expires_at < datetime('now', '-30 days')")
|
||||
return err
|
||||
}
|
||||
|
||||
// ── Internal ────────────────────────────────
|
||||
|
||||
func (s *UserStore) getBy(ctx context.Context, col, val string) (*models.User, error) {
|
||||
func (s *UserStore) scanOne(ctx context.Context, query string, args ...interface{}) (*models.User, error) {
|
||||
var u models.User
|
||||
var displayName, avatarURL sql.NullString
|
||||
var settingsJSON []byte
|
||||
err := DB.QueryRowContext(ctx, fmt.Sprintf(`
|
||||
SELECT id, username, email, password_hash, display_name, avatar_url,
|
||||
role, is_active, settings, created_at, updated_at, last_login_at
|
||||
FROM users WHERE %s = ?`, col), val).Scan(
|
||||
&u.ID, &u.Username, &u.Email, &u.PasswordHash, &displayName, &avatarURL,
|
||||
&u.Role, &u.IsActive, &settingsJSON, st(&u.CreatedAt), st(&u.UpdatedAt), stN(&u.LastLoginAt),
|
||||
var dn, av, ph sql.NullString
|
||||
var extID, hdl sql.NullString
|
||||
var sj []byte
|
||||
err := DB.QueryRowContext(ctx, query, args...).Scan(
|
||||
&u.ID, &u.Username, &u.Email, &ph, &dn, &av,
|
||||
&u.Role, &u.IsActive, &sj, st(&u.CreatedAt), st(&u.UpdatedAt), stN(&u.LastLoginAt),
|
||||
&u.AuthSource, &extID, &hdl,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
u.DisplayName = NullableString(displayName)
|
||||
u.AvatarURL = NullableString(avatarURL)
|
||||
ScanJSON(settingsJSON, &u.Settings)
|
||||
u.PasswordHash = NullableString(ph)
|
||||
u.DisplayName = NullableString(dn)
|
||||
u.AvatarURL = NullableString(av)
|
||||
if extID.Valid { u.ExternalID = &extID.String }
|
||||
u.Handle = NullableString(hdl)
|
||||
ScanJSON(sj, &u.Settings)
|
||||
return &u, nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user