feat(user): customizable ban duration with automatic lift (#2945)

Admins can now set an optional ban expiry when banning a user; the ban
lifts automatically once it passes. Empty expiry keeps the permanent
behavior.

- schema: User.ban_expires (nillable, additive)
- inventory: LiftExpiredBan restores expired bans; Upsert propagates
  ban_expires for banned statuses and clears it otherwise
- login, password reset, and SSO all lift expired bans before the ban
  check
- admin user dialog: ban-expiry datetime field shown while a banned
  status is selected
- tests: lift matrix (expired/future/permanent/active) and upsert
  propagation

Authored By: TDvorak <info@tdvorak.dev>

Generated with [Devin](https://devin.ai)

Co-Authored-By: Devin <158243242+devin-ai-integration[bot]@users.noreply.github.com>
pull/3582/head
Tomas Dvorak 2 weeks ago
parent b85216fd51
commit e1c2864ae7

File diff suppressed because one or more lines are too long

@ -482,6 +482,7 @@ var (
{Name: "nick", Type: field.TypeString, Size: 100},
{Name: "password", Type: field.TypeString, Nullable: true},
{Name: "status", Type: field.TypeEnum, Enums: []string{"active", "inactive", "manual_banned", "sys_banned"}, Default: "active"},
{Name: "ban_expires", Type: field.TypeTime, Nullable: true},
{Name: "storage", Type: field.TypeInt64, Default: 0},
{Name: "two_factor_secret", Type: field.TypeString, Nullable: true},
{Name: "avatar", Type: field.TypeString, Nullable: true},
@ -496,7 +497,7 @@ var (
ForeignKeys: []*schema.ForeignKey{
{
Symbol: "users_groups_users",
Columns: []*schema.Column{UsersColumns[12]},
Columns: []*schema.Column{UsersColumns[13]},
RefColumns: []*schema.Column{GroupsColumns[0]},
OnDelete: schema.NoAction,
},

@ -15118,6 +15118,7 @@ type UserMutation struct {
nick *string
password *string
status *user.Status
ban_expires *time.Time
storage *int64
addstorage *int64
two_factor_secret *string
@ -15531,6 +15532,55 @@ func (m *UserMutation) ResetStatus() {
m.status = nil
}
// SetBanExpires sets the "ban_expires" field.
func (m *UserMutation) SetBanExpires(t time.Time) {
m.ban_expires = &t
}
// BanExpires returns the value of the "ban_expires" field in the mutation.
func (m *UserMutation) BanExpires() (r time.Time, exists bool) {
v := m.ban_expires
if v == nil {
return
}
return *v, true
}
// OldBanExpires returns the old "ban_expires" field's value of the User entity.
// If the User object wasn't provided to the builder, the object is fetched from the database.
// An error is returned if the mutation operation is not UpdateOne, or the database query fails.
func (m *UserMutation) OldBanExpires(ctx context.Context) (v *time.Time, err error) {
if !m.op.Is(OpUpdateOne) {
return v, errors.New("OldBanExpires is only allowed on UpdateOne operations")
}
if m.id == nil || m.oldValue == nil {
return v, errors.New("OldBanExpires requires an ID field in the mutation")
}
oldValue, err := m.oldValue(ctx)
if err != nil {
return v, fmt.Errorf("querying old value for OldBanExpires: %w", err)
}
return oldValue.BanExpires, nil
}
// ClearBanExpires clears the value of the "ban_expires" field.
func (m *UserMutation) ClearBanExpires() {
m.ban_expires = nil
m.clearedFields[user.FieldBanExpires] = struct{}{}
}
// BanExpiresCleared returns if the "ban_expires" field was cleared in this mutation.
func (m *UserMutation) BanExpiresCleared() bool {
_, ok := m.clearedFields[user.FieldBanExpires]
return ok
}
// ResetBanExpires resets all changes to the "ban_expires" field.
func (m *UserMutation) ResetBanExpires() {
m.ban_expires = nil
delete(m.clearedFields, user.FieldBanExpires)
}
// SetStorage sets the "storage" field.
func (m *UserMutation) SetStorage(i int64) {
m.storage = &i
@ -16276,7 +16326,7 @@ func (m *UserMutation) Type() string {
// order to get all numeric fields that were incremented/decremented, call
// AddedFields().
func (m *UserMutation) Fields() []string {
fields := make([]string, 0, 12)
fields := make([]string, 0, 13)
if m.created_at != nil {
fields = append(fields, user.FieldCreatedAt)
}
@ -16298,6 +16348,9 @@ func (m *UserMutation) Fields() []string {
if m.status != nil {
fields = append(fields, user.FieldStatus)
}
if m.ban_expires != nil {
fields = append(fields, user.FieldBanExpires)
}
if m.storage != nil {
fields = append(fields, user.FieldStorage)
}
@ -16335,6 +16388,8 @@ func (m *UserMutation) Field(name string) (ent.Value, bool) {
return m.Password()
case user.FieldStatus:
return m.Status()
case user.FieldBanExpires:
return m.BanExpires()
case user.FieldStorage:
return m.Storage()
case user.FieldTwoFactorSecret:
@ -16368,6 +16423,8 @@ func (m *UserMutation) OldField(ctx context.Context, name string) (ent.Value, er
return m.OldPassword(ctx)
case user.FieldStatus:
return m.OldStatus(ctx)
case user.FieldBanExpires:
return m.OldBanExpires(ctx)
case user.FieldStorage:
return m.OldStorage(ctx)
case user.FieldTwoFactorSecret:
@ -16436,6 +16493,13 @@ func (m *UserMutation) SetField(name string, value ent.Value) error {
}
m.SetStatus(v)
return nil
case user.FieldBanExpires:
v, ok := value.(time.Time)
if !ok {
return fmt.Errorf("unexpected type %T for field %s", value, name)
}
m.SetBanExpires(v)
return nil
case user.FieldStorage:
v, ok := value.(int64)
if !ok {
@ -16522,6 +16586,9 @@ func (m *UserMutation) ClearedFields() []string {
if m.FieldCleared(user.FieldPassword) {
fields = append(fields, user.FieldPassword)
}
if m.FieldCleared(user.FieldBanExpires) {
fields = append(fields, user.FieldBanExpires)
}
if m.FieldCleared(user.FieldTwoFactorSecret) {
fields = append(fields, user.FieldTwoFactorSecret)
}
@ -16551,6 +16618,9 @@ func (m *UserMutation) ClearField(name string) error {
case user.FieldPassword:
m.ClearPassword()
return nil
case user.FieldBanExpires:
m.ClearBanExpires()
return nil
case user.FieldTwoFactorSecret:
m.ClearTwoFactorSecret()
return nil
@ -16589,6 +16659,9 @@ func (m *UserMutation) ResetField(name string) error {
case user.FieldStatus:
m.ResetStatus()
return nil
case user.FieldBanExpires:
m.ResetBanExpires()
return nil
case user.FieldStorage:
m.ResetStorage()
return nil

@ -415,11 +415,11 @@ func init() {
// user.NickValidator is a validator for the "nick" field. It is called by the builders before save.
user.NickValidator = userDescNick.Validators[0].(func(string) error)
// userDescStorage is the schema descriptor for storage field.
userDescStorage := userFields[4].Descriptor()
userDescStorage := userFields[5].Descriptor()
// user.DefaultStorage holds the default value on creation for the storage field.
user.DefaultStorage = userDescStorage.Default.(int64)
// userDescSettings is the schema descriptor for settings field.
userDescSettings := userFields[7].Descriptor()
userDescSettings := userFields[8].Descriptor()
// user.DefaultSettings holds the default value on creation for the settings field.
user.DefaultSettings = userDescSettings.Default.(*types.UserSetting)
}

@ -25,6 +25,10 @@ func (User) Fields() []ent.Field {
field.Enum("status").
Values("active", "inactive", "manual_banned", "sys_banned").
Default("active"),
// ban_expires lifts a banned status after this time; nil = permanent.
field.Time("ban_expires").
Optional().
Nillable(),
field.Int64("storage").
Default(0),
field.String("two_factor_secret").

@ -34,6 +34,8 @@ type User struct {
Password string `json:"-"`
// Status holds the value of the "status" field.
Status user.Status `json:"status,omitempty"`
// BanExpires holds the value of the "ban_expires" field.
BanExpires *time.Time `json:"ban_expires,omitempty"`
// Storage holds the value of the "storage" field.
Storage int64 `json:"storage,omitempty"`
// TwoFactorSecret holds the value of the "two_factor_secret" field.
@ -171,7 +173,7 @@ func (*User) scanValues(columns []string) ([]any, error) {
values[i] = new(sql.NullInt64)
case user.FieldEmail, user.FieldNick, user.FieldPassword, user.FieldStatus, user.FieldTwoFactorSecret, user.FieldAvatar:
values[i] = new(sql.NullString)
case user.FieldCreatedAt, user.FieldUpdatedAt, user.FieldDeletedAt:
case user.FieldCreatedAt, user.FieldUpdatedAt, user.FieldDeletedAt, user.FieldBanExpires:
values[i] = new(sql.NullTime)
default:
values[i] = new(sql.UnknownType)
@ -237,6 +239,13 @@ func (u *User) assignValues(columns []string, values []any) error {
} else if value.Valid {
u.Status = user.Status(value.String)
}
case user.FieldBanExpires:
if value, ok := values[i].(*sql.NullTime); !ok {
return fmt.Errorf("unexpected type %T for field ban_expires", values[i])
} else if value.Valid {
u.BanExpires = new(time.Time)
*u.BanExpires = value.Time
}
case user.FieldStorage:
if value, ok := values[i].(*sql.NullInt64); !ok {
return fmt.Errorf("unexpected type %T for field storage", values[i])
@ -372,6 +381,11 @@ func (u *User) String() string {
builder.WriteString("status=")
builder.WriteString(fmt.Sprintf("%v", u.Status))
builder.WriteString(", ")
if v := u.BanExpires; v != nil {
builder.WriteString("ban_expires=")
builder.WriteString(v.Format(time.ANSIC))
}
builder.WriteString(", ")
builder.WriteString("storage=")
builder.WriteString(fmt.Sprintf("%v", u.Storage))
builder.WriteString(", ")

@ -31,6 +31,8 @@ const (
FieldPassword = "password"
// FieldStatus holds the string denoting the status field in the database.
FieldStatus = "status"
// FieldBanExpires holds the string denoting the ban_expires field in the database.
FieldBanExpires = "ban_expires"
// FieldStorage holds the string denoting the storage field in the database.
FieldStorage = "storage"
// FieldTwoFactorSecret holds the string denoting the two_factor_secret field in the database.
@ -136,6 +138,7 @@ var Columns = []string{
FieldNick,
FieldPassword,
FieldStatus,
FieldBanExpires,
FieldStorage,
FieldTwoFactorSecret,
FieldAvatar,
@ -248,6 +251,11 @@ func ByStatus(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldStatus, opts...).ToFunc()
}
// ByBanExpires orders the results by the ban_expires field.
func ByBanExpires(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldBanExpires, opts...).ToFunc()
}
// ByStorage orders the results by the storage field.
func ByStorage(opts ...sql.OrderTermOption) OrderOption {
return sql.OrderByField(FieldStorage, opts...).ToFunc()

@ -85,6 +85,11 @@ func Password(v string) predicate.User {
return predicate.User(sql.FieldEQ(FieldPassword, v))
}
// BanExpires applies equality check predicate on the "ban_expires" field. It's identical to BanExpiresEQ.
func BanExpires(v time.Time) predicate.User {
return predicate.User(sql.FieldEQ(FieldBanExpires, v))
}
// Storage applies equality check predicate on the "storage" field. It's identical to StorageEQ.
func Storage(v int64) predicate.User {
return predicate.User(sql.FieldEQ(FieldStorage, v))
@ -460,6 +465,56 @@ func StatusNotIn(vs ...Status) predicate.User {
return predicate.User(sql.FieldNotIn(FieldStatus, vs...))
}
// BanExpiresEQ applies the EQ predicate on the "ban_expires" field.
func BanExpiresEQ(v time.Time) predicate.User {
return predicate.User(sql.FieldEQ(FieldBanExpires, v))
}
// BanExpiresNEQ applies the NEQ predicate on the "ban_expires" field.
func BanExpiresNEQ(v time.Time) predicate.User {
return predicate.User(sql.FieldNEQ(FieldBanExpires, v))
}
// BanExpiresIn applies the In predicate on the "ban_expires" field.
func BanExpiresIn(vs ...time.Time) predicate.User {
return predicate.User(sql.FieldIn(FieldBanExpires, vs...))
}
// BanExpiresNotIn applies the NotIn predicate on the "ban_expires" field.
func BanExpiresNotIn(vs ...time.Time) predicate.User {
return predicate.User(sql.FieldNotIn(FieldBanExpires, vs...))
}
// BanExpiresGT applies the GT predicate on the "ban_expires" field.
func BanExpiresGT(v time.Time) predicate.User {
return predicate.User(sql.FieldGT(FieldBanExpires, v))
}
// BanExpiresGTE applies the GTE predicate on the "ban_expires" field.
func BanExpiresGTE(v time.Time) predicate.User {
return predicate.User(sql.FieldGTE(FieldBanExpires, v))
}
// BanExpiresLT applies the LT predicate on the "ban_expires" field.
func BanExpiresLT(v time.Time) predicate.User {
return predicate.User(sql.FieldLT(FieldBanExpires, v))
}
// BanExpiresLTE applies the LTE predicate on the "ban_expires" field.
func BanExpiresLTE(v time.Time) predicate.User {
return predicate.User(sql.FieldLTE(FieldBanExpires, v))
}
// BanExpiresIsNil applies the IsNil predicate on the "ban_expires" field.
func BanExpiresIsNil() predicate.User {
return predicate.User(sql.FieldIsNull(FieldBanExpires))
}
// BanExpiresNotNil applies the NotNil predicate on the "ban_expires" field.
func BanExpiresNotNil() predicate.User {
return predicate.User(sql.FieldNotNull(FieldBanExpires))
}
// StorageEQ applies the EQ predicate on the "storage" field.
func StorageEQ(v int64) predicate.User {
return predicate.User(sql.FieldEQ(FieldStorage, v))

@ -114,6 +114,20 @@ func (uc *UserCreate) SetNillableStatus(u *user.Status) *UserCreate {
return uc
}
// SetBanExpires sets the "ban_expires" field.
func (uc *UserCreate) SetBanExpires(t time.Time) *UserCreate {
uc.mutation.SetBanExpires(t)
return uc
}
// SetNillableBanExpires sets the "ban_expires" field if the given value is not nil.
func (uc *UserCreate) SetNillableBanExpires(t *time.Time) *UserCreate {
if t != nil {
uc.SetBanExpires(*t)
}
return uc
}
// SetStorage sets the "storage" field.
func (uc *UserCreate) SetStorage(i int64) *UserCreate {
uc.mutation.SetStorage(i)
@ -468,6 +482,10 @@ func (uc *UserCreate) createSpec() (*User, *sqlgraph.CreateSpec) {
_spec.SetField(user.FieldStatus, field.TypeEnum, value)
_node.Status = value
}
if value, ok := uc.mutation.BanExpires(); ok {
_spec.SetField(user.FieldBanExpires, field.TypeTime, value)
_node.BanExpires = &value
}
if value, ok := uc.mutation.Storage(); ok {
_spec.SetField(user.FieldStorage, field.TypeInt64, value)
_node.Storage = value
@ -765,6 +783,24 @@ func (u *UserUpsert) UpdateStatus() *UserUpsert {
return u
}
// SetBanExpires sets the "ban_expires" field.
func (u *UserUpsert) SetBanExpires(v time.Time) *UserUpsert {
u.Set(user.FieldBanExpires, v)
return u
}
// UpdateBanExpires sets the "ban_expires" field to the value that was provided on create.
func (u *UserUpsert) UpdateBanExpires() *UserUpsert {
u.SetExcluded(user.FieldBanExpires)
return u
}
// ClearBanExpires clears the value of the "ban_expires" field.
func (u *UserUpsert) ClearBanExpires() *UserUpsert {
u.SetNull(user.FieldBanExpires)
return u
}
// SetStorage sets the "storage" field.
func (u *UserUpsert) SetStorage(v int64) *UserUpsert {
u.Set(user.FieldStorage, v)
@ -992,6 +1028,27 @@ func (u *UserUpsertOne) UpdateStatus() *UserUpsertOne {
})
}
// SetBanExpires sets the "ban_expires" field.
func (u *UserUpsertOne) SetBanExpires(v time.Time) *UserUpsertOne {
return u.Update(func(s *UserUpsert) {
s.SetBanExpires(v)
})
}
// UpdateBanExpires sets the "ban_expires" field to the value that was provided on create.
func (u *UserUpsertOne) UpdateBanExpires() *UserUpsertOne {
return u.Update(func(s *UserUpsert) {
s.UpdateBanExpires()
})
}
// ClearBanExpires clears the value of the "ban_expires" field.
func (u *UserUpsertOne) ClearBanExpires() *UserUpsertOne {
return u.Update(func(s *UserUpsert) {
s.ClearBanExpires()
})
}
// SetStorage sets the "storage" field.
func (u *UserUpsertOne) SetStorage(v int64) *UserUpsertOne {
return u.Update(func(s *UserUpsert) {
@ -1404,6 +1461,27 @@ func (u *UserUpsertBulk) UpdateStatus() *UserUpsertBulk {
})
}
// SetBanExpires sets the "ban_expires" field.
func (u *UserUpsertBulk) SetBanExpires(v time.Time) *UserUpsertBulk {
return u.Update(func(s *UserUpsert) {
s.SetBanExpires(v)
})
}
// UpdateBanExpires sets the "ban_expires" field to the value that was provided on create.
func (u *UserUpsertBulk) UpdateBanExpires() *UserUpsertBulk {
return u.Update(func(s *UserUpsert) {
s.UpdateBanExpires()
})
}
// ClearBanExpires clears the value of the "ban_expires" field.
func (u *UserUpsertBulk) ClearBanExpires() *UserUpsertBulk {
return u.Update(func(s *UserUpsert) {
s.ClearBanExpires()
})
}
// SetStorage sets the "storage" field.
func (u *UserUpsertBulk) SetStorage(v int64) *UserUpsertBulk {
return u.Update(func(s *UserUpsert) {

@ -126,6 +126,26 @@ func (uu *UserUpdate) SetNillableStatus(u *user.Status) *UserUpdate {
return uu
}
// SetBanExpires sets the "ban_expires" field.
func (uu *UserUpdate) SetBanExpires(t time.Time) *UserUpdate {
uu.mutation.SetBanExpires(t)
return uu
}
// SetNillableBanExpires sets the "ban_expires" field if the given value is not nil.
func (uu *UserUpdate) SetNillableBanExpires(t *time.Time) *UserUpdate {
if t != nil {
uu.SetBanExpires(*t)
}
return uu
}
// ClearBanExpires clears the value of the "ban_expires" field.
func (uu *UserUpdate) ClearBanExpires() *UserUpdate {
uu.mutation.ClearBanExpires()
return uu
}
// SetStorage sets the "storage" field.
func (uu *UserUpdate) SetStorage(i int64) *UserUpdate {
uu.mutation.ResetStorage()
@ -624,6 +644,12 @@ func (uu *UserUpdate) sqlSave(ctx context.Context) (n int, err error) {
if value, ok := uu.mutation.Status(); ok {
_spec.SetField(user.FieldStatus, field.TypeEnum, value)
}
if value, ok := uu.mutation.BanExpires(); ok {
_spec.SetField(user.FieldBanExpires, field.TypeTime, value)
}
if uu.mutation.BanExpiresCleared() {
_spec.ClearField(user.FieldBanExpires, field.TypeTime)
}
if value, ok := uu.mutation.Storage(); ok {
_spec.SetField(user.FieldStorage, field.TypeInt64, value)
}
@ -1145,6 +1171,26 @@ func (uuo *UserUpdateOne) SetNillableStatus(u *user.Status) *UserUpdateOne {
return uuo
}
// SetBanExpires sets the "ban_expires" field.
func (uuo *UserUpdateOne) SetBanExpires(t time.Time) *UserUpdateOne {
uuo.mutation.SetBanExpires(t)
return uuo
}
// SetNillableBanExpires sets the "ban_expires" field if the given value is not nil.
func (uuo *UserUpdateOne) SetNillableBanExpires(t *time.Time) *UserUpdateOne {
if t != nil {
uuo.SetBanExpires(*t)
}
return uuo
}
// ClearBanExpires clears the value of the "ban_expires" field.
func (uuo *UserUpdateOne) ClearBanExpires() *UserUpdateOne {
uuo.mutation.ClearBanExpires()
return uuo
}
// SetStorage sets the "storage" field.
func (uuo *UserUpdateOne) SetStorage(i int64) *UserUpdateOne {
uuo.mutation.ResetStorage()
@ -1673,6 +1719,12 @@ func (uuo *UserUpdateOne) sqlSave(ctx context.Context) (_node *User, err error)
if value, ok := uuo.mutation.Status(); ok {
_spec.SetField(user.FieldStatus, field.TypeEnum, value)
}
if value, ok := uuo.mutation.BanExpires(); ok {
_spec.SetField(user.FieldBanExpires, field.TypeTime, value)
}
if uuo.mutation.BanExpiresCleared() {
_spec.ClearField(user.FieldBanExpires, field.TypeTime)
}
if value, ok := uuo.mutation.Storage(); ok {
_spec.SetField(user.FieldStorage, field.TypeInt64, value)
}

@ -1384,7 +1384,9 @@
"deleteXUsers": "Delete {{num}} users",
"confirmBatchDelete": "Are you sure you want to delete {{num}} users?",
"calibrateStorage": "Calibrate storage",
"calibrateStorageSuccess": "Storage calibrated successfully."
"calibrateStorageSuccess": "Storage calibrated successfully.",
"banExpires": "Ban expires",
"banExpiresDes": "Optional. The ban lifts automatically after this time; empty means permanent."
},
"file": {
"deleteXFiles": "Delete {{num}} files",

@ -1384,7 +1384,9 @@
"deleteXUsers": "删除 {{num}} 个用户",
"confirmBatchDelete": "确认要删除 {{num}} 个用户?",
"calibrateStorage": "校准存储空间",
"calibrateStorageSuccess": "存储空间校准成功"
"calibrateStorageSuccess": "存储空间校准成功",
"banExpires": "封禁截止时间",
"banExpiresDes": "可选。到期后自动解除封禁;留空表示永久封禁。"
},
"file": {
"deleteXFiles": "删除 {{num}} 个文件",

@ -257,6 +257,7 @@ export interface User extends CommonMixin {
avatar?: string;
credit?: number;
group_expires?: string;
ban_expires?: string;
notify_date?: string;
group_users?: number;
previous_group?: number;

@ -65,6 +65,30 @@ const UserForm = ({ reload, setLoading }: { reload: () => void; setLoading: (loa
[setUser],
);
const banned = values.status == UserStatus.manual_banned || values.status == UserStatus.sys_banned;
// datetime-local inputs want "YYYY-MM-DDTHH:mm" in local time; the API
// stores RFC3339.
const toLocalInput = (iso?: string) => {
if (!iso) {
return "";
}
const d = new Date(iso);
if (isNaN(d.getTime())) {
return "";
}
const pad = (n: number) => String(n).padStart(2, "0");
return `${d.getFullYear()}-${pad(d.getMonth() + 1)}-${pad(d.getDate())}T${pad(d.getHours())}:${pad(d.getMinutes())}`;
};
const onBanExpiresChange = useCallback(
(e: React.ChangeEvent<HTMLInputElement>) => {
const v = e.target.value;
setUser((prev) => ({ ...prev, ban_expires: v ? new Date(v).toISOString() : undefined }));
},
[setUser],
);
const onGroupChange = useCallback(
(value: string) => {
setUser((prev) => ({ ...prev, group_users: parseInt(value) }));
@ -179,6 +203,18 @@ const UserForm = ({ reload, setLoading }: { reload: () => void; setLoading: (loa
</DenseSelect>
</FormControl>
</SettingForm>
{banned && (
<SettingForm title={t("user.banExpires")} noContainer lgWidth={6}>
<DenseFilledTextField
fullWidth
type="datetime-local"
value={toLocalInput(values.ban_expires)}
onChange={onBanExpiresChange}
slotProps={{ inputLabel: { shrink: true } }}
/>
<NoMarginHelperText>{t("user.banExpiresDes")}</NoMarginHelperText>
</SettingForm>
)}
<SettingForm title={t("user.group")} noContainer lgWidth={6}>
<GroupSelectionInput value={values.group_users?.toString() ?? ""} onChange={onGroupChange} fullWidth />
</SettingForm>

@ -0,0 +1,93 @@
package inventory
import (
"context"
"testing"
"time"
"github.com/cloudreve/Cloudreve/v4/ent/enttest"
entuser "github.com/cloudreve/Cloudreve/v4/ent/user"
"github.com/cloudreve/Cloudreve/v4/pkg/boolset"
"github.com/stretchr/testify/require"
)
func TestLiftExpiredBan(t *testing.T) {
client := enttest.Open(t, "sqlite3", "file:"+t.Name()+"?mode=memory&cache=shared")
t.Cleanup(func() { require.NoError(t, client.Close()) })
ctx := context.Background()
group := client.Group.Create().SetName("g").SetPermissions(&boolset.BooleanSet{}).SaveX(ctx)
past := time.Now().Add(-time.Hour)
future := time.Now().Add(time.Hour)
uc := NewUserClient(client)
t.Run("expired manual ban lifted", func(t *testing.T) {
u := client.User.Create().SetEmail("lift-exp@example.com").SetNick("u").
SetStatus(entuser.StatusManualBanned).SetBanExpires(past).SetGroup(group).SaveX(ctx)
out, err := uc.LiftExpiredBan(ctx, u)
require.NoError(t, err)
require.Equal(t, entuser.StatusActive, out.Status)
require.Nil(t, client.User.GetX(ctx, u.ID).BanExpires)
})
t.Run("expired sys ban lifted", func(t *testing.T) {
u := client.User.Create().SetEmail("lift-sys@example.com").SetNick("u").
SetStatus(entuser.StatusSysBanned).SetBanExpires(past).SetGroup(group).SaveX(ctx)
out, err := uc.LiftExpiredBan(ctx, u)
require.NoError(t, err)
require.Equal(t, entuser.StatusActive, out.Status)
})
t.Run("future ban untouched", func(t *testing.T) {
u := client.User.Create().SetEmail("lift-future@example.com").SetNick("u").
SetStatus(entuser.StatusManualBanned).SetBanExpires(future).SetGroup(group).SaveX(ctx)
out, err := uc.LiftExpiredBan(ctx, u)
require.NoError(t, err)
require.Equal(t, entuser.StatusManualBanned, out.Status)
require.NotNil(t, client.User.GetX(ctx, u.ID).BanExpires)
})
t.Run("permanent ban untouched", func(t *testing.T) {
u := client.User.Create().SetEmail("lift-perm@example.com").SetNick("u").
SetStatus(entuser.StatusManualBanned).SetGroup(group).SaveX(ctx)
out, err := uc.LiftExpiredBan(ctx, u)
require.NoError(t, err)
require.Equal(t, entuser.StatusManualBanned, out.Status)
})
t.Run("active user untouched", func(t *testing.T) {
u := client.User.Create().SetEmail("lift-active@example.com").SetNick("u").
SetStatus(entuser.StatusActive).SetGroup(group).SaveX(ctx)
out, err := uc.LiftExpiredBan(ctx, u)
require.NoError(t, err)
require.Equal(t, entuser.StatusActive, out.Status)
})
}
func TestUpsertBanExpires(t *testing.T) {
client := enttest.Open(t, "sqlite3", "file:"+t.Name()+"?mode=memory&cache=shared")
t.Cleanup(func() { require.NoError(t, client.Close()) })
ctx := context.Background()
group := client.Group.Create().SetName("g").SetPermissions(&boolset.BooleanSet{}).SaveX(ctx)
uc := NewUserClient(client)
future := time.Now().Add(24 * time.Hour).UTC().Truncate(time.Second)
u := client.User.Create().SetEmail("upsert-ban@example.com").SetNick("u").
SetStatus(entuser.StatusActive).SetGroup(group).SaveX(ctx)
// Ban with expiry
u.Status = entuser.StatusManualBanned
u.BanExpires = &future
_, err := uc.Upsert(ctx, u, "", "")
require.NoError(t, err)
got := client.User.GetX(ctx, u.ID)
require.Equal(t, future, got.BanExpires.UTC())
// Unban clears the expiry
u.Status = entuser.StatusActive
u.BanExpires = nil
_, err = uc.Upsert(ctx, u, "", "")
require.NoError(t, err)
require.Nil(t, client.User.GetX(ctx, u.ID).BanExpires)
}

@ -61,6 +61,9 @@ type (
GetActiveByID(ctx context.Context, id int) (*ent.User, error)
// SetStatus Set user to given status
SetStatus(ctx context.Context, u *ent.User, status user.Status) (*ent.User, error)
// LiftExpiredBan restores a banned user whose ban_expires has passed.
// It returns the (possibly updated) user; permanent bans are untouched.
LiftExpiredBan(ctx context.Context, u *ent.User) (*ent.User, error)
// AnonymousUser returns the anonymous user.
AnonymousUser(ctx context.Context) (*ent.User, error)
// GetLoginUserByID returns the login user by its ID. It emits some errors and fallback to anonymous user.
@ -364,6 +367,24 @@ func (c *userClient) SetStatus(ctx context.Context, u *ent.User, status user.Sta
return c.client.User.UpdateOne(u).SetStatus(status).Save(ctx)
}
func (c *userClient) LiftExpiredBan(ctx context.Context, u *ent.User) (*ent.User, error) {
banned := u.Status == user.StatusManualBanned || u.Status == user.StatusSysBanned
if !banned || u.BanExpires == nil || u.BanExpires.After(time.Now()) {
return u, nil
}
if err := c.client.User.UpdateOneID(u.ID).
SetStatus(user.StatusActive).
ClearBanExpires().
Exec(ctx); err != nil {
return nil, err
}
u.Status = user.StatusActive
u.BanExpires = nil
return u, nil
}
func (c *userClient) Create(ctx context.Context, args *NewUserArgs) (*ent.User, error) {
// Try to check if there's user with same email.
if existedUser, err := c.GetByEmail(ctx, args.Email); err == nil {
@ -579,6 +600,12 @@ func (c *userClient) Upsert(ctx context.Context, u *ent.User, password, twoFa st
SetStatus(u.Status).
SetGroupID(u.GroupUsers)
if u.Status == user.StatusManualBanned || u.Status == user.StatusSysBanned {
q.SetNillableBanExpires(u.BanExpires)
} else {
q.ClearBanExpires()
}
if password != "" {
pwdDigest, err := digestPassword(password)
if err != nil {

@ -88,6 +88,10 @@ func (service *UserResetEmailService) Reset(c *gin.Context) error {
return serializer.NewError(serializer.CodeUserNotFound, "User not found", err)
}
if u, err = userClient.LiftExpiredBan(c, u); err != nil {
return serializer.NewError(serializer.CodeDBError, "Failed to lift expired ban", err)
}
if u.Status == user.StatusManualBanned || u.Status == user.StatusSysBanned {
return serializer.NewError(serializer.CodeUserBaned, "This user is banned", nil)
}
@ -127,6 +131,9 @@ func (service *UserLoginService) Login(c *gin.Context) (*ent.User, string, error
ctx := context.WithValue(c, inventory.LoadUserGroup{}, true)
expectedUser, err := userClient.GetByEmail(ctx, service.UserName)
if err == nil {
expectedUser, err = userClient.LiftExpiredBan(ctx, expectedUser)
}
// 一系列校验
if err != nil {

@ -258,6 +258,9 @@ func (service *SSOExchangeService) SSOExchange(c *gin.Context) (any, error) {
if err != nil {
return nil, serializer.NewError(serializer.CodeUserNotFound, "User not found", err)
}
if u, err = dep.UserClient().LiftExpiredBan(c, u); err != nil {
return nil, serializer.NewError(serializer.CodeDBError, "Failed to lift expired ban", err)
}
if err := checkUserStatus(u); err != nil {
return nil, err
}
@ -272,6 +275,9 @@ func ssoResolveUser(c *gin.Context, dep dependency.Dep, sso *setting.SSO, email,
u, err := userClient.GetByEmail(c, email)
if err == nil {
if u, err = userClient.LiftExpiredBan(c, u); err != nil {
return nil, serializer.NewError(serializer.CodeDBError, "Failed to lift expired ban", err)
}
if err := checkUserStatus(u); err != nil {
return nil, err
}

Loading…
Cancel
Save