diff --git a/internal/handler/admin.go b/internal/handler/admin.go index 58729de..3dee9d5 100644 --- a/internal/handler/admin.go +++ b/internal/handler/admin.go @@ -42,7 +42,7 @@ func createUserHandler(svc *service.Container) gin.HandlerFunc { } u, _, err := svc.Auth.Register(c.Request.Context(), req.Username, req.Password) if err != nil { - writeUserMutationError(c, err) + writeUserMutationError(c, svc, err) return } // Admin-created users are intentionally normal viewers by default. @@ -93,7 +93,7 @@ func updateUserHandler(svc *service.Container) gin.HandlerFunc { c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) return } else if existing != nil && existing.ID != userID { - writeUserMutationError(c, service.ErrUsernameTaken) + writeUserMutationError(c, svc, service.ErrUsernameTaken) return } updates := map[string]any{"username": nextUsername} @@ -167,12 +167,13 @@ func annotateProtectedUsers(ctx context.Context, svc *service.Container, users [ return nil } -func writeUserMutationError(c *gin.Context, err error) { +func writeUserMutationError(c *gin.Context, svc *service.Container, err error) { switch { case errors.Is(err, service.ErrUsernameTaken): c.JSON(http.StatusConflict, gin.H{"error": "username already taken"}) case errors.Is(err, service.ErrUserLimitReached): - c.JSON(http.StatusBadRequest, gin.H{"error": "user limit reached: maximum 20 users"}) + maxUsers := service.LicensedMaxUsers(c.Request.Context(), svc.Repo) + c.JSON(http.StatusBadRequest, gin.H{"error": "user limit reached", "max_users": maxUsers}) default: c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) } diff --git a/internal/handler/license.go b/internal/handler/license.go index e78eb56..11d3978 100644 --- a/internal/handler/license.go +++ b/internal/handler/license.go @@ -36,22 +36,26 @@ type licenseActivateReq struct { } type licenseServerSignedResp struct { - Valid bool `json:"valid"` - LicenseType string `json:"license_type"` - ExpiryDate *string `json:"expiry_date"` - MaxDevices int `json:"max_devices"` - DaysRemaining *int `json:"days_remaining"` - NextHeartbeat string `json:"next_heartbeat"` - Signature string `json:"signature"` + Valid bool `json:"valid"` + LicenseType string `json:"license_type"` + ExpiryDate *string `json:"expiry_date"` + MaxDevices int `json:"max_devices"` + MaxUsers *int `json:"max_users"` + DaysRemaining *int `json:"days_remaining"` + NextHeartbeat string `json:"next_heartbeat"` + Signature string `json:"signature"` + LegacySignature bool `json:"-"` } type licenseServerStatusResp struct { - Valid bool `json:"valid"` - LicenseType *string `json:"license_type"` - ExpiryDate *string `json:"expiry_date"` - DaysRemaining *int `json:"days_remaining"` - DeviceName string `json:"device_name"` - IsActive bool `json:"is_active"` + Valid bool `json:"valid"` + LicenseType *string `json:"license_type"` + ExpiryDate *string `json:"expiry_date"` + MaxUsers *int `json:"max_users"` + UnlimitedUsers bool `json:"unlimited_users"` + DaysRemaining *int `json:"days_remaining"` + DeviceName string `json:"device_name"` + IsActive bool `json:"is_active"` } func licenseActivateHandler(svc *service.Container) gin.HandlerFunc { @@ -88,7 +92,7 @@ func licenseActivateHandler(svc *service.Container) gin.HandlerFunc { c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()}) return } - if err := client.verifySigned(upstream); err != nil { + if err := client.verifySigned(&upstream); err != nil { c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()}) return } @@ -117,6 +121,8 @@ func licenseStatusHandler(svc *service.Container) gin.HandlerFunc { if upstream.ExpiryDate != nil { state.ExpiryDate = *upstream.ExpiryDate } + state.MaxUsers = upstream.MaxUsers + state.UnlimitedUsers = upstream.UnlimitedUsers state.DaysRemaining = upstream.DaysRemaining if upstream.DeviceName != "" { state.DeviceName = upstream.DeviceName @@ -132,10 +138,11 @@ func licenseStatusHandler(svc *service.Container) gin.HandlerFunc { } active := state.Valid && !licenseStateExpired(state.ExpiryDate) c.JSON(http.StatusOK, gin.H{ - "active": active, - "message": licenseStatusMessage(active, err), - "max_users": service.LicensedMaxUsers(c.Request.Context(), svc.Repo), - "activation": licenseActivationView(state), + "active": active, + "message": licenseStatusMessage(active, err), + "max_users": licenseStatusMaxUsers(state), + "unlimited_users": state.Valid && !licenseStateExpired(state.ExpiryDate) && state.UnlimitedUsers, + "activation": licenseActivationView(state), }) } } @@ -160,7 +167,7 @@ func licenseHeartbeatHandler(svc *service.Container) gin.HandlerFunc { c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()}) return } - if err := client.verifySigned(upstream); err != nil { + if err := client.verifySigned(&upstream); err != nil { c.JSON(http.StatusBadGateway, gin.H{"error": err.Error()}) return } @@ -248,10 +255,45 @@ func (c *licenseClient) do(req *http.Request, out any) error { return json.Unmarshal(data, out) } -func (c *licenseClient) verifySigned(resp licenseServerSignedResp) error { +func (c *licenseClient) verifySigned(resp *licenseServerSignedResp) error { if c.hmacSecret == "" { return nil } + unsigned := struct { + Valid bool `json:"valid"` + LicenseType string `json:"license_type"` + ExpiryDate *string `json:"expiry_date"` + MaxDevices int `json:"max_devices"` + MaxUsers *int `json:"max_users"` + DaysRemaining *int `json:"days_remaining"` + NextHeartbeat string `json:"next_heartbeat"` + }{ + Valid: resp.Valid, + LicenseType: resp.LicenseType, + ExpiryDate: resp.ExpiryDate, + MaxDevices: resp.MaxDevices, + MaxUsers: resp.MaxUsers, + DaysRemaining: resp.DaysRemaining, + NextHeartbeat: resp.NextHeartbeat, + } + payload, err := json.Marshal(unsigned) + if err != nil { + return err + } + mac := hmac.New(sha256.New, []byte(c.hmacSecret)) + _, _ = mac.Write(payload) + expected := hex.EncodeToString(mac.Sum(nil)) + if !hmac.Equal([]byte(expected), []byte(resp.Signature)) { + if c.verifyLegacySigned(*resp) { + resp.LegacySignature = true + return nil + } + return errors.New("license server signature verification failed") + } + return nil +} + +func (c *licenseClient) verifyLegacySigned(resp licenseServerSignedResp) bool { unsigned := struct { Valid bool `json:"valid"` LicenseType string `json:"license_type"` @@ -269,15 +311,12 @@ func (c *licenseClient) verifySigned(resp licenseServerSignedResp) error { } payload, err := json.Marshal(unsigned) if err != nil { - return err + return false } mac := hmac.New(sha256.New, []byte(c.hmacSecret)) _, _ = mac.Write(payload) expected := hex.EncodeToString(mac.Sum(nil)) - if !hmac.Equal([]byte(expected), []byte(resp.Signature)) { - return errors.New("license server signature verification failed") - } - return nil + return hmac.Equal([]byte(expected), []byte(resp.Signature)) } func ensureLicenseDeviceID(ctx context.Context, svc *service.Container, candidate string) (string, error) { @@ -313,18 +352,34 @@ func licenseStateFromSigned(resp licenseServerSignedResp, deviceID, deviceName s expiry = *resp.ExpiryDate } return service.LicenseActivationState{ - Valid: resp.Valid, - LicenseType: resp.LicenseType, - ExpiryDate: expiry, - MaxDevices: resp.MaxDevices, - DaysRemaining: resp.DaysRemaining, - NextHeartbeat: resp.NextHeartbeat, - DeviceID: deviceID, - DeviceName: deviceName, - UpdatedAt: time.Now().Format(time.RFC3339), + Valid: resp.Valid, + LicenseType: resp.LicenseType, + ExpiryDate: expiry, + MaxDevices: resp.MaxDevices, + MaxUsers: resp.MaxUsers, + UnlimitedUsers: !resp.LegacySignature && resp.MaxUsers == nil, + DaysRemaining: resp.DaysRemaining, + NextHeartbeat: resp.NextHeartbeat, + DeviceID: deviceID, + DeviceName: deviceName, + UpdatedAt: time.Now().Format(time.RFC3339), } } +func licenseStatusMaxUsers(state service.LicenseActivationState) any { + active := state.Valid && !licenseStateExpired(state.ExpiryDate) + if active { + if state.UnlimitedUsers { + return nil + } + if state.MaxUsers != nil && *state.MaxUsers > 0 { + return *state.MaxUsers + } + return service.LicensedUserLimit + } + return service.OpenSourceUserLimit +} + func persistLicenseState(ctx context.Context, svc *service.Container, state service.LicenseActivationState) error { data, err := json.Marshal(state) if err != nil { @@ -357,6 +412,8 @@ func licenseActivationView(state service.LicenseActivationState) gin.H { "device_name": state.DeviceName, "plan": state.LicenseType, "max_activations": state.MaxDevices, + "max_users": state.MaxUsers, + "unlimited_users": state.UnlimitedUsers, "expires_at": emptyAsNil(state.ExpiryDate), "valid": state.Valid && !licenseStateExpired(state.ExpiryDate), "heartbeat_at": updatedAt, diff --git a/internal/handler/license_test.go b/internal/handler/license_test.go new file mode 100644 index 0000000..408e77e --- /dev/null +++ b/internal/handler/license_test.go @@ -0,0 +1,32 @@ +package handler + +import ( + "testing" + + "github.com/ShukeBta/MediaStationGo/internal/service" +) + +func TestLicenseStatusMaxUsersUsesLicensedLimit(t *testing.T) { + maxUsers := 25 + state := service.LicenseActivationState{Valid: true, MaxUsers: &maxUsers} + + if got := licenseStatusMaxUsers(state); got != maxUsers { + t.Fatalf("expected licensed max users %d, got %#v", maxUsers, got) + } +} + +func TestLicenseStatusMaxUsersAllowsUnlimited(t *testing.T) { + state := service.LicenseActivationState{Valid: true, UnlimitedUsers: true} + + if got := licenseStatusMaxUsers(state); got != nil { + t.Fatalf("expected unlimited max users to be nil, got %#v", got) + } +} + +func TestLicenseStatusMaxUsersFallsBackToOpenSourceLimit(t *testing.T) { + state := service.LicenseActivationState{} + + if got := licenseStatusMaxUsers(state); got != service.OpenSourceUserLimit { + t.Fatalf("expected open-source max users %d, got %#v", service.OpenSourceUserLimit, got) + } +} diff --git a/internal/service/auth_user_limits_test.go b/internal/service/auth_user_limits_test.go index 36e8d29..4570e1b 100644 --- a/internal/service/auth_user_limits_test.go +++ b/internal/service/auth_user_limits_test.go @@ -2,6 +2,7 @@ package service import ( "context" + "encoding/json" "errors" "fmt" "testing" @@ -23,7 +24,7 @@ func newAuthTestServices(t *testing.T) (*repository.Container, *AuthService, *Pr if err != nil { t.Fatal(err) } - if err := db.AutoMigrate(&model.User{}, &model.UserPermission{}, &model.RefreshToken{}, &model.TelegramBinding{}); err != nil { + if err := db.AutoMigrate(&model.User{}, &model.UserPermission{}, &model.RefreshToken{}, &model.TelegramBinding{}, &model.Setting{}); err != nil { t.Fatal(err) } repos := repository.New(db) @@ -57,6 +58,44 @@ func TestRegisterRejectsMoreThanTwentyUsers(t *testing.T) { } } +func TestRegisterUsesLicensedUserLimit(t *testing.T) { + ctx := context.Background() + repos, auth, _, _ := newAuthTestServices(t) + maxUsers := 25 + state := LicenseActivationState{Valid: true, LicenseType: "plus", MaxUsers: &maxUsers} + raw, _ := json.Marshal(state) + if err := repos.Setting.Set(ctx, LicenseSettingActivation, string(raw)); err != nil { + t.Fatal(err) + } + for i := 0; i < OpenSourceUserLimit; i++ { + if err := repos.User.Create(ctx, &model.User{ + Username: fmt.Sprintf("licensed-%02d", i), + PasswordHash: "hash", + Role: "user", + Tier: "free", + }); err != nil { + t.Fatal(err) + } + } + + if _, _, err := auth.Register(ctx, "extra", "password"); err != nil { + t.Fatalf("licensed user limit should allow user 21: %v", err) + } +} + +func TestLicensedMaxUsersCanBeUnlimited(t *testing.T) { + ctx := context.Background() + repos, _, _, _ := newAuthTestServices(t) + state := LicenseActivationState{Valid: true, LicenseType: "enterprise", UnlimitedUsers: true} + raw, _ := json.Marshal(state) + if err := repos.Setting.Set(ctx, LicenseSettingActivation, string(raw)); err != nil { + t.Fatal(err) + } + if got := LicensedMaxUsers(ctx, repos); got <= 1_000_000 { + t.Fatalf("unlimited license should return a very high limit, got %d", got) + } +} + func TestRegisterDefaultsAdultLibrariesHidden(t *testing.T) { _, auth, _, _ := newAuthTestServices(t) user, _, err := auth.Register(context.Background(), "viewer", "password") diff --git a/internal/service/license.go b/internal/service/license.go index 6a56902..94c0003 100644 --- a/internal/service/license.go +++ b/internal/service/license.go @@ -3,6 +3,7 @@ package service import ( "context" "encoding/json" + "math" "time" "github.com/ShukeBta/MediaStationGo/internal/repository" @@ -16,19 +17,28 @@ const ( ) type LicenseActivationState struct { - Valid bool `json:"valid"` - LicenseType string `json:"license_type,omitempty"` - ExpiryDate string `json:"expiry_date,omitempty"` - MaxDevices int `json:"max_devices,omitempty"` - DaysRemaining *int `json:"days_remaining,omitempty"` - NextHeartbeat string `json:"next_heartbeat,omitempty"` - DeviceID string `json:"device_id,omitempty"` - DeviceName string `json:"device_name,omitempty"` - UpdatedAt string `json:"updated_at,omitempty"` + Valid bool `json:"valid"` + LicenseType string `json:"license_type,omitempty"` + ExpiryDate string `json:"expiry_date,omitempty"` + MaxDevices int `json:"max_devices,omitempty"` + MaxUsers *int `json:"max_users,omitempty"` + UnlimitedUsers bool `json:"unlimited_users,omitempty"` + DaysRemaining *int `json:"days_remaining,omitempty"` + NextHeartbeat string `json:"next_heartbeat,omitempty"` + DeviceID string `json:"device_id,omitempty"` + DeviceName string `json:"device_name,omitempty"` + UpdatedAt string `json:"updated_at,omitempty"` } func LicensedMaxUsers(ctx context.Context, repos *repository.Container) int64 { - if LicenseActive(ctx, repos) { + state, ok := loadLicenseActivationState(ctx, repos) + if ok && state.Valid && !licenseExpired(state.ExpiryDate) { + if state.UnlimitedUsers { + return math.MaxInt64 + } + if state.MaxUsers != nil && *state.MaxUsers > 0 { + return int64(*state.MaxUsers) + } return LicensedUserLimit } return OpenSourceUserLimit @@ -38,15 +48,23 @@ func LicenseActive(ctx context.Context, repos *repository.Container) bool { if repos == nil || repos.Setting == nil { return false } + state, ok := loadLicenseActivationState(ctx, repos) + return ok && state.Valid && !licenseExpired(state.ExpiryDate) +} + +func loadLicenseActivationState(ctx context.Context, repos *repository.Container) (LicenseActivationState, bool) { + if repos == nil || repos.Setting == nil { + return LicenseActivationState{}, false + } raw, err := repos.Setting.Get(ctx, LicenseSettingActivation) if err != nil || raw == "" { - return false + return LicenseActivationState{}, false } var state LicenseActivationState if err := json.Unmarshal([]byte(raw), &state); err != nil { - return false + return LicenseActivationState{}, false } - return state.Valid && !licenseExpired(state.ExpiryDate) + return state, true } func licenseExpired(expiry string) bool { diff --git a/web/src/api/license.ts b/web/src/api/license.ts index 76e87d5..ce0f2a4 100644 --- a/web/src/api/license.ts +++ b/web/src/api/license.ts @@ -12,6 +12,8 @@ export interface LicenseActivation { device_name?: string plan?: string max_activations?: number + max_users?: number | null + unlimited_users?: boolean /** ISO8601 — null means perpetual */ expires_at?: string | null valid: boolean @@ -25,7 +27,8 @@ export interface LicenseStatus { /** Whether a license is currently active */ active: boolean activation?: LicenseActivation - max_users?: number + max_users?: number | null + unlimited_users?: boolean /** Error or status message */ message?: string } diff --git a/web/src/pages/AdminPage.tsx b/web/src/pages/AdminPage.tsx index db076b1..ec4ef40 100644 --- a/web/src/pages/AdminPage.tsx +++ b/web/src/pages/AdminPage.tsx @@ -4,6 +4,7 @@ import { KeyRound, Pencil, Plus, ShieldCheck, Trash2, X } from 'lucide-react' import { adminAPI } from '../api/admin' import { libraryAPI } from '../api/library' +import { licenseAPI, type LicenseStatus } from '../api/license' import type { Library, User } from '../types' import { APIConfigsPanel } from '../components/APIConfigsPanel' import { ManagementShortcuts } from '../components/ManagementShortcuts' @@ -163,15 +164,30 @@ function LibraryPanel() { function UsersPanel() { const [users, setUsers] = useState([]) + const [licenseStatus, setLicenseStatus] = useState(null) const [username, setUsername] = useState('') const [password, setPassword] = useState('') const [editingID, setEditingID] = useState(null) const [editingUsername, setEditingUsername] = useState('') - const refresh = () => adminAPI.listUsers().then(setUsers) + const refresh = async () => { + const [nextUsers, nextLicense] = await Promise.all([ + adminAPI.listUsers(), + licenseAPI.status().catch(() => null), + ]) + setUsers(nextUsers) + setLicenseStatus(nextLicense) + } useEffect(() => { refresh().catch(() => undefined) }, []) + const unlimitedUsers = + licenseStatus?.active === true && + (licenseStatus.unlimited_users === true || licenseStatus.max_users == null) + const maxUsers = unlimitedUsers ? null : (licenseStatus?.max_users ?? 20) + const userLimitReached = maxUsers != null && users.length >= maxUsers + const userLimitLabel = unlimitedUsers ? '不限制' : String(maxUsers) + const handleCreate = async (e: FormEvent) => { e.preventDefault() try { @@ -182,7 +198,7 @@ function UsersPanel() { await refresh() } catch (err: unknown) { const msg = - (err as { response?: { data?: { error?: string } } })?.response?.data?.error ?? + userCreateErrorMessage(err) ?? '添加用户失败' toast.error(msg) } @@ -236,7 +252,7 @@ function UsersPanel() {

用户管理

- 已创建 {users.length}/20 个用户;新增用户默认只有媒体库浏览、播放、外部播放器与第三方客户端观看权限。 + 已创建 {users.length}/{userLimitLabel} 个用户;新增用户默认只有媒体库浏览、播放、外部播放器与第三方客户端观看权限。

@@ -249,7 +265,7 @@ function UsersPanel() { placeholder="用户名" value={username} onChange={(e) => setUsername(e.target.value)} - disabled={users.length >= 20} + disabled={userLimitReached} /> setPassword(e.target.value)} - disabled={users.length >= 20} + disabled={userLimitReached} /> - @@ -356,3 +372,12 @@ function UsersPanel() { ) } + +function userCreateErrorMessage(err: unknown): string | undefined { + const data = (err as { response?: { data?: { error?: string; max_users?: number } } })?.response?.data + if (!data?.error) return undefined + if (data.error === 'user limit reached' && data.max_users != null) { + return `用户数量已达到授权上限:${data.max_users} 人` + } + return data.error +} diff --git a/web/src/pages/LicensePage.tsx b/web/src/pages/LicensePage.tsx index 209f2a4..8f2cb91 100644 --- a/web/src/pages/LicensePage.tsx +++ b/web/src/pages/LicensePage.tsx @@ -27,6 +27,11 @@ function fmtDateTime(iso: string | null | undefined): string { return new Date(iso).toLocaleString('zh-CN') } +function fmtUserLimit(maxUsers: number | null | undefined, unlimited?: boolean): string { + if (unlimited || maxUsers == null) return '不限制' + return `${maxUsers} 人` +} + // ── Page ── export function LicensePage() { @@ -198,7 +203,10 @@ export function LicensePage() { /> - + )}