package repo import ( "context" "crypto/rand" "database/sql" "encoding/hex" "encoding/json" "errors" "fmt" "strings" "time" "github.com/agentmail/gateway/internal/db" "github.com/agentmail/gateway/internal/models" "github.com/google/uuid" "golang.org/x/crypto/bcrypt" ) const ( bcryptCost = 12 sessionTTL = 7 * 24 * time.Hour userSelectCols = `user_id, username, display_name, password_hash, role, status, created_at, last_login, allowed_agents, allowed_paths` ) var ( ErrUserNotFound = errors.New("user not found") ErrBadCredentials = errors.New("invalid username or password") ErrUserDisabled = errors.New("user disabled") ErrNameTaken = errors.New("name already taken by an agent or user") ErrSessionInvalid = errors.New("session invalid or expired") ErrInvalidUsername = errors.New("username must be 2-64 chars of [a-z0-9._-]") ErrAlreadySetup = errors.New("system already initialized") ) // ---------- 命名空间校验 ---------- // 三维地址的 name 位由人类用户与 Agent 共用,因此必须全局唯一 func nameTaken(ctx context.Context, name string) (bool, error) { var n int err := db.DB.QueryRowContext(ctx, ` SELECT (SELECT COUNT(*) FROM users WHERE username = $1) + (SELECT COUNT(*) FROM agents WHERE agent_name = $1) `, name).Scan(&n) return n > 0, err } // AgentNameAvailable 供 Agent 注册前校验(不与人类用户重名) func AgentNameAvailable(ctx context.Context, agentName string) (bool, error) { var n int err := db.DB.QueryRowContext(ctx, `SELECT COUNT(*) FROM users WHERE username = $1`, agentName).Scan(&n) return n == 0, err } func validUsername(name string) bool { if len(name) < 2 || len(name) > 64 { return false } for _, r := range name { ok := (r >= 'a' && r <= 'z') || (r >= '0' && r <= '9') || r == '.' || r == '_' || r == '-' if !ok { return false } } // 保留字:human 是兼容别名,不能被真实用户占用 return name != "human" } // ---------- User CRUD ---------- func scanUser(row *sql.Row) (*models.User, error) { var u models.User var agentsJSON, pathsJSON []byte err := row.Scan(&u.ID, &u.Username, &u.DisplayName, &u.PasswordHash, &u.Role, &u.Status, &u.CreatedAt, &u.LastLogin, &agentsJSON, &pathsJSON) if err != nil { if errors.Is(err, sql.ErrNoRows) { return nil, ErrUserNotFound } return nil, err } u.AllowedAgents = decodeStrList(agentsJSON) u.AllowedPaths = decodeStrList(pathsJSON) return &u, nil } func decodeStrList(raw []byte) []string { out := []string{} if len(raw) > 0 { _ = json.Unmarshal(raw, &out) } if out == nil { out = []string{} } return out } func CreateUser(ctx context.Context, username, password, displayName, role string, allowedAgents, allowedPaths []string) (*models.User, error) { username = strings.ToLower(strings.TrimSpace(username)) if !validUsername(username) { return nil, ErrInvalidUsername } if role != "admin" { role = "user" } if displayName == "" { displayName = username } taken, err := nameTaken(ctx, username) if err != nil { return nil, err } if taken { return nil, ErrNameTaken } hash, err := bcrypt.GenerateFromPassword([]byte(password), bcryptCost) if err != nil { return nil, err } agentsJSON, _ := json.Marshal(normalizeList(allowedAgents)) pathsJSON, _ := json.Marshal(normalizeList(allowedPaths)) row := db.DB.QueryRowContext(ctx, ` INSERT INTO users (username, display_name, password_hash, role, allowed_agents, allowed_paths) VALUES ($1, $2, $3, $4, $5, $6) RETURNING `+userSelectCols, username, displayName, string(hash), role, agentsJSON, pathsJSON) return scanUser(row) } func GetUserByName(ctx context.Context, username string) (*models.User, error) { return scanUser(db.DB.QueryRowContext(ctx, `SELECT `+userSelectCols+` FROM users WHERE username = $1`, strings.ToLower(strings.TrimSpace(username)))) } func GetUserByID(ctx context.Context, id uuid.UUID) (*models.User, error) { return scanUser(db.DB.QueryRowContext(ctx, `SELECT `+userSelectCols+` FROM users WHERE user_id = $1`, id)) } func ListUsers(ctx context.Context) ([]models.User, error) { rows, err := db.DB.QueryContext(ctx, `SELECT `+userSelectCols+` FROM users ORDER BY created_at ASC`) if err != nil { return nil, err } defer rows.Close() users := []models.User{} for rows.Next() { var u models.User var agentsJSON, pathsJSON []byte if err := rows.Scan(&u.ID, &u.Username, &u.DisplayName, &u.PasswordHash, &u.Role, &u.Status, &u.CreatedAt, &u.LastLogin, &agentsJSON, &pathsJSON); err != nil { return nil, err } u.AllowedAgents = decodeStrList(agentsJSON) u.AllowedPaths = decodeStrList(pathsJSON) users = append(users, u) } return users, nil } // UserUpdate 描述一次用户更新;nil 字段表示不改 type UserUpdate struct { DisplayName *string Role *string Status *string AllowedAgents *[]string AllowedPaths *[]string } func UpdateUser(ctx context.Context, id uuid.UUID, up UserUpdate) (*models.User, error) { var agentsJSON, pathsJSON *string if up.AllowedAgents != nil { b, _ := json.Marshal(normalizeList(*up.AllowedAgents)) s := string(b) agentsJSON = &s } if up.AllowedPaths != nil { b, _ := json.Marshal(normalizeList(*up.AllowedPaths)) s := string(b) pathsJSON = &s } row := db.DB.QueryRowContext(ctx, ` UPDATE users SET display_name = COALESCE($2, display_name), role = COALESCE($3, role), status = COALESCE($4, status), allowed_agents = COALESCE($5`+db.JSONCast()+`, allowed_agents), allowed_paths = COALESCE($6`+db.JSONCast()+`, allowed_paths) WHERE user_id = $1 RETURNING `+userSelectCols, id, up.DisplayName, up.Role, up.Status, agentsJSON, pathsJSON) return scanUser(row) } // normalizeList 去空白、去空项、去重,保持顺序 func normalizeList(in []string) []string { out := make([]string, 0, len(in)) seen := map[string]bool{} for _, s := range in { s = strings.TrimSpace(s) if s == "" || seen[s] { continue } seen[s] = true out = append(out, s) } return out } func SetPassword(ctx context.Context, id uuid.UUID, newPassword string) error { hash, err := bcrypt.GenerateFromPassword([]byte(newPassword), bcryptCost) if err != nil { return err } tag, err := db.DB.ExecContext(ctx, `UPDATE users SET password_hash = $2 WHERE user_id = $1`, id, string(hash)) if err != nil { return err } if n, _ := tag.RowsAffected(); n == 0 { return ErrUserNotFound } // 改密后踢掉该用户所有会话 _, _ = db.DB.ExecContext(ctx, `DELETE FROM user_sessions WHERE user_id = $1`, id) return nil } func DisableUser(ctx context.Context, id uuid.UUID) error { tag, err := db.DB.ExecContext(ctx, `UPDATE users SET status = 'disabled' WHERE user_id = $1`, id) if err != nil { return err } if n, _ := tag.RowsAffected(); n == 0 { return ErrUserNotFound } _, _ = db.DB.ExecContext(ctx, `DELETE FROM user_sessions WHERE user_id = $1`, id) return nil } func CountAdmins(ctx context.Context) (int, error) { var n int err := db.DB.QueryRowContext(ctx, `SELECT COUNT(*) FROM users WHERE role = 'admin' AND status = 'active'`).Scan(&n) return n, err } // EnsureAdminUser 首次启动时创建默认管理员(幂等) func EnsureAdminUser(ctx context.Context, username, password string) (*models.User, bool, error) { if n, err := CountAdmins(ctx); err != nil { return nil, false, err } else if n > 0 { u, err := GetUserByName(ctx, username) if err != nil && !errors.Is(err, ErrUserNotFound) { return nil, false, err } return u, false, nil } u, err := CreateUser(ctx, username, password, "管理员", "admin", nil, nil) if err != nil { return nil, false, err } return u, true, nil } // ---------- 登录 / 会话令牌 ---------- func Authenticate(ctx context.Context, username, password string) (*models.User, error) { u, err := GetUserByName(ctx, username) if err != nil { if errors.Is(err, ErrUserNotFound) { // 统一错误,避免暴露用户是否存在 return nil, ErrBadCredentials } return nil, err } if u.Status != "active" { return nil, ErrUserDisabled } if bcrypt.CompareHashAndPassword([]byte(u.PasswordHash), []byte(password)) != nil { return nil, ErrBadCredentials } _, _ = db.DB.ExecContext(ctx, `UPDATE users SET last_login = NOW() WHERE user_id = $1`, u.ID) return u, nil } func newToken() (string, error) { b := make([]byte, 32) if _, err := rand.Read(b); err != nil { return "", err } return hex.EncodeToString(b), nil } func CreateUserSession(ctx context.Context, userID uuid.UUID, userAgent string) (string, time.Time, error) { token, err := newToken() if err != nil { return "", time.Time{}, err } expires := time.Now().Add(sessionTTL) if len(userAgent) > 256 { userAgent = userAgent[:256] } _, err = db.DB.ExecContext(ctx, ` INSERT INTO user_sessions (token, user_id, expires_at, user_agent) VALUES ($1, $2, $3, $4)`, token, userID, expires, userAgent) if err != nil { return "", time.Time{}, err } // 顺手清理过期令牌 _, _ = db.DB.ExecContext(ctx, `DELETE FROM user_sessions WHERE expires_at < NOW()`) return token, expires, nil } // ResolveUserSession 校验令牌并滑动续期 func ResolveUserSession(ctx context.Context, token string) (*models.User, error) { if token == "" { return nil, ErrSessionInvalid } var userID uuid.UUID err := db.DB.QueryRowContext(ctx, ` SELECT user_id FROM user_sessions WHERE token = $1 AND expires_at > NOW()`, token).Scan(&userID) if err != nil { if errors.Is(err, sql.ErrNoRows) { return nil, ErrSessionInvalid } return nil, err } u, err := GetUserByID(ctx, userID) if err != nil { return nil, err } if u.Status != "active" { return nil, ErrUserDisabled } _, _ = db.DB.ExecContext(ctx, `UPDATE user_sessions SET expires_at = $2 WHERE token = $1`, token, time.Now().Add(sessionTTL)) return u, nil } func DeleteUserSession(ctx context.Context, token string) error { _, err := db.DB.ExecContext(ctx, `DELETE FROM user_sessions WHERE token = $1`, token) return err } // ---------- 人类用户候选(供地址补全) ---------- func ListActiveUsernames(ctx context.Context) ([]string, error) { rows, err := db.DB.QueryContext(ctx, `SELECT username FROM users WHERE status = 'active' ORDER BY username`) if err != nil { return nil, err } defer rows.Close() out := []string{} for rows.Next() { var s string if err := rows.Scan(&s); err == nil { out = append(out, s) } } return out, rows.Err() } // ---------- 会话归属 ---------- func SetSessionOwner(ctx context.Context, sessionID, userID uuid.UUID) error { _, err := db.DB.ExecContext(ctx, `UPDATE sessions SET owner_user_id = $2 WHERE session_id = $1 AND owner_user_id IS NULL`, sessionID, userID) return err } // SessionOwnerUsername 返回会话归属人类用户名;无归属时返回空串 func SessionOwnerUsername(ctx context.Context, sessionID uuid.UUID) (string, error) { var name *string err := db.DB.QueryRowContext(ctx, ` SELECT u.username FROM sessions s LEFT JOIN users u ON u.user_id = s.owner_user_id WHERE s.session_id = $1`, sessionID).Scan(&name) if err != nil { if errors.Is(err, sql.ErrNoRows) { return "", fmt.Errorf("session %s not found", sessionID) } return "", err } if name == nil { return "", nil } return *name, nil } // UserCanAccessSession 判断用户能否访问该会话:owner、或在邮件收发/抄送中出现,或 admin func UserCanAccessSession(ctx context.Context, u *models.User, sessionID uuid.UUID) (bool, error) { if u.IsAdmin() { return true, nil } var n int err := db.DB.QueryRowContext(ctx, ` SELECT COUNT(*) FROM sessions s WHERE s.session_id = $1 AND (s.owner_user_id = $2 OR EXISTS ( SELECT 1 FROM mails m WHERE m.session_id = s.session_id AND (m.from_name = $3 OR m.to_name = $3 OR `+db.CCHas("m.cc_list", 3)+`) )) `, sessionID, u.ID, u.Username).Scan(&n) return n > 0, err } // FirstAdminUsername 返回最早创建的可用管理员用户名(用于无归属会话的兜底决策人) func FirstAdminUsername(ctx context.Context) (string, error) { var name string err := db.DB.QueryRowContext(ctx, ` SELECT username FROM users WHERE role = 'admin' AND status = 'active' ORDER BY created_at ASC LIMIT 1`).Scan(&name) if err != nil { if errors.Is(err, sql.ErrNoRows) { return "", nil } return "", err } return name, nil } // RandomPassword 生成一个随机初始密码(首次启动无 ADMIN_PASSWORD 时使用) func RandomPassword(n int) string { const charset = "abcdefghijkmnopqrstuvwxyzABCDEFGHJKLMNPQRSTUVWXYZ23456789" b := make([]byte, n) if _, err := rand.Read(b); err != nil { return "ChangeMe" + fmt.Sprint(time.Now().Unix()) } for i := range b { b[i] = charset[int(b[i])%len(charset)] } return string(b) } // ---------- Setup(首次初始化管理员) ---------- // NeedsSetup 返回系统是否尚未初始化(没有任何用户) func NeedsSetup(ctx context.Context) (bool, error) { var n int err := db.DB.QueryRowContext(ctx, `SELECT COUNT(*) FROM users`).Scan(&n) return n == 0, err } // SetupFirstAdmin 在系统尚无任何用户时创建首个管理员。 // 已初始化时返回 ErrAlreadySetup,避免被用作后门。 func SetupFirstAdmin(ctx context.Context, username, password, displayName string) (*models.User, error) { empty, err := NeedsSetup(ctx) if err != nil { return nil, err } if !empty { return nil, ErrAlreadySetup } if displayName == "" { displayName = username } return CreateUser(ctx, username, password, displayName, "admin", nil, nil) } // ---------- 可选目录候选(供权限设置界面) ---------- // AllWorkspaceNames 汇总所有 Agent 注册过的工作区名,供管理员挑选可访问目录 func AllWorkspaceNames(ctx context.Context) ([]string, error) { rows, err := db.DB.QueryContext(ctx, ` SELECT DISTINCT ws->>'name' AS name FROM agents, jsonb_array_elements(workspaces) AS ws WHERE COALESCE(ws->>'name', '') <> '' ORDER BY name`) if err != nil { return []string{}, err } defer rows.Close() out := []string{} for rows.Next() { var s string if err := rows.Scan(&s); err == nil { out = append(out, s) } } return out, rows.Err() } // IsHumanUser 判断某个三维地址 name 位是否为人类用户 func IsHumanUser(ctx context.Context, name string) (bool, error) { var n int err := db.DB.QueryRowContext(ctx, `SELECT COUNT(*) FROM users WHERE username = $1`, name).Scan(&n) return n > 0, err }