fix: honor licensed user limits

This commit is contained in:
ShukeBta
2026-05-30 21:48:09 +08:00
parent e08b827861
commit de3f1f3bb4
8 changed files with 243 additions and 60 deletions
+5 -4
View File
@@ -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()})
}
+91 -34
View File
@@ -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,
+32
View File
@@ -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)
}
}
+40 -1
View File
@@ -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")
+31 -13
View File
@@ -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 {