fix: report licensed instances on startup

This commit is contained in:
ShukeBta
2026-06-25 13:46:02 +08:00
parent 74aa654c92
commit a4a93b7fed
7 changed files with 190 additions and 26 deletions
+47 -15
View File
@@ -114,7 +114,8 @@ func licenseStatusHandler(svc *service.Container) gin.HandlerFunc {
return func(c *gin.Context) {
state, _ := loadLicenseState(c.Request.Context(), svc)
client, err := newLicenseClient(c.Request.Context(), svc)
if err == nil {
hasLicenseKey := strings.TrimSpace(state.LicenseKey) != ""
if err == nil && hasLicenseKey {
deviceID, idErr := ensureLicenseDeviceID(c.Request.Context(), svc, state.DeviceID)
if idErr == nil {
deviceName, _ := svc.Repo.Setting.Get(c.Request.Context(), licenseDeviceNameSetting)
@@ -142,6 +143,17 @@ func licenseStatusHandler(svc *service.Container) gin.HandlerFunc {
}
}
}
} else if err == nil && state.Valid {
deviceID, idErr := ensureLicenseDeviceID(c.Request.Context(), svc, state.DeviceID)
if idErr == nil {
if refreshed, ok, _, getErr := refreshLicenseServerStatus(c.Request.Context(), client, state, deviceID); getErr == nil && ok {
state = refreshed
_ = persistLicenseState(c.Request.Context(), svc, state)
} else if getErr == nil {
state.Valid = false
_ = persistLicenseState(c.Request.Context(), svc, state)
}
}
}
active := state.Valid && !licenseStateExpired(state.ExpiryDate)
c.JSON(http.StatusOK, gin.H{
@@ -179,17 +191,13 @@ func RunLicenseHeartbeatLoop(ctx context.Context, svc *service.Container) {
if svc == nil {
return
}
run := func() {
state, sent, err := maybeSendLicenseHeartbeat(ctx, svc, licenseHeartbeatInterval)
if err != nil {
if svc.Log != nil {
svc.Log.Warn("license heartbeat failed", zap.Error(err))
}
return
}
if sent && svc.Log != nil {
svc.Log.Info("license heartbeat sent", zap.String("device_id", state.DeviceID))
}
run := func(interval time.Duration) {
state, sent, err := maybeSendLicenseHeartbeat(ctx, svc, interval)
logLicenseHeartbeatResult(svc, state, sent, err)
}
runStartup := func() {
state, sent, err := maybeSendStartupLicenseHeartbeat(ctx, svc)
logLicenseHeartbeatResult(svc, state, sent, err)
}
timer := time.NewTimer(licenseHeartbeatStartupDelay)
@@ -198,7 +206,7 @@ func RunLicenseHeartbeatLoop(ctx context.Context, svc *service.Container) {
case <-ctx.Done():
return
case <-timer.C:
run()
runStartup()
}
ticker := time.NewTicker(licenseHeartbeatCheckInterval)
@@ -208,11 +216,35 @@ func RunLicenseHeartbeatLoop(ctx context.Context, svc *service.Container) {
case <-ctx.Done():
return
case <-ticker.C:
run()
run(licenseHeartbeatInterval)
}
}
}
func logLicenseHeartbeatResult(svc *service.Container, state service.LicenseActivationState, sent bool, err error) {
if svc == nil || svc.Log == nil {
return
}
if err != nil {
svc.Log.Warn("license heartbeat failed", zap.Error(err))
return
}
if sent {
svc.Log.Info("license heartbeat sent", zap.String("device_id", state.DeviceID))
}
}
func maybeSendStartupLicenseHeartbeat(ctx context.Context, svc *service.Container) (service.LicenseActivationState, bool, error) {
state, err := loadLicenseState(ctx, svc)
if err != nil {
return state, false, nil
}
if strings.TrimSpace(state.LicenseKey) == "" {
return state, false, nil
}
return maybeSendLicenseHeartbeat(ctx, svc, 0)
}
func maybeSendLicenseHeartbeat(ctx context.Context, svc *service.Container, interval time.Duration) (service.LicenseActivationState, bool, error) {
state, err := loadLicenseState(ctx, svc)
if err != nil {
@@ -247,7 +279,7 @@ func maybeSendLicenseHeartbeat(ctx context.Context, svc *service.Container, inte
}
func licenseHeartbeatEligible(state service.LicenseActivationState) bool {
return strings.TrimSpace(state.LicenseKey) != "" || state.Valid
return strings.TrimSpace(state.LicenseKey) != ""
}
func licenseHeartbeatDue(state service.LicenseActivationState, interval time.Duration) bool {
+124 -2
View File
@@ -157,6 +157,102 @@ func TestLicenseHeartbeatDueUsesTwelveHourWindow(t *testing.T) {
}
}
func TestStartupLicenseHeartbeatIgnoresTwelveHourWindow(t *testing.T) {
heartbeatCount := 0
maxUsers := 40
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/v1/heartbeat":
heartbeatCount++
resp := licenseServerSignedResp{
Valid: true,
LicenseType: "subscription",
MaxDevices: 2,
MaxUsers: &maxUsers,
NextHeartbeat: time.Now().Add(time.Hour).Format(time.RFC3339),
}
resp.Signature = signLicenseTestPayload("test-secret", resp)
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(resp)
case "/api/v1/status/device-1":
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{
"valid": true,
"license_type": "subscription",
"max_devices": 2,
"max_users": 40,
"unlimited_users": false,
"device_name": "NAS",
"is_active": true
}`))
default:
t.Fatalf("unexpected upstream path %s", r.URL.Path)
}
}))
defer upstream.Close()
svc := newLicenseHandlerTestService(t)
if err := svc.Repo.Setting.Set(t.Context(), licenseServerURLSetting, upstream.URL); err != nil {
t.Fatal(err)
}
if err := svc.Repo.Setting.Set(t.Context(), licenseHMACSecretSetting, "test-secret"); err != nil {
t.Fatal(err)
}
state := service.LicenseActivationState{
Valid: true,
LicenseKey: "MS-ABCD-EFGH-JKLM-NPQR",
DeviceID: "device-1",
DeviceName: "NAS",
UpdatedAt: time.Now().Format(time.RFC3339),
}
if err := persistLicenseState(t.Context(), svc, state); err != nil {
t.Fatal(err)
}
refreshed, sent, err := maybeSendStartupLicenseHeartbeat(t.Context(), svc)
if err != nil {
t.Fatalf("startup heartbeat: %v", err)
}
if !sent || heartbeatCount != 1 {
t.Fatalf("startup heartbeat should be sent once, sent=%v count=%d", sent, heartbeatCount)
}
if refreshed.MaxUsers == nil || *refreshed.MaxUsers != 40 {
t.Fatalf("startup heartbeat should refresh licensed user capacity, got %+v", refreshed)
}
}
func TestStartupLicenseHeartbeatSkipsStateWithoutStoredKey(t *testing.T) {
heartbeatCount := 0
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
heartbeatCount++
w.WriteHeader(http.StatusInternalServerError)
}))
defer upstream.Close()
svc := newLicenseHandlerTestService(t)
if err := svc.Repo.Setting.Set(t.Context(), licenseServerURLSetting, upstream.URL); err != nil {
t.Fatal(err)
}
if err := svc.Repo.Setting.Set(t.Context(), licenseHMACSecretSetting, "test-secret"); err != nil {
t.Fatal(err)
}
if err := persistLicenseState(t.Context(), svc, service.LicenseActivationState{
Valid: true,
DeviceID: "device-1",
UpdatedAt: time.Now().Format(time.RFC3339),
}); err != nil {
t.Fatal(err)
}
_, sent, err := maybeSendStartupLicenseHeartbeat(t.Context(), svc)
if err != nil {
t.Fatalf("startup heartbeat should skip without error, got %v", err)
}
if sent || heartbeatCount != 0 {
t.Fatalf("startup heartbeat without stored license key should be skipped, sent=%v count=%d", sent, heartbeatCount)
}
}
func TestLicenseHeartbeatEligibleRequiresActivationState(t *testing.T) {
if licenseHeartbeatEligible(service.LicenseActivationState{DeviceID: "device-only"}) {
t.Fatalf("device id alone should not trigger automatic license heartbeat")
@@ -164,8 +260,34 @@ func TestLicenseHeartbeatEligibleRequiresActivationState(t *testing.T) {
if !licenseHeartbeatEligible(service.LicenseActivationState{LicenseKey: "MS-KEY"}) {
t.Fatalf("stored license key should trigger automatic license heartbeat")
}
if !licenseHeartbeatEligible(service.LicenseActivationState{Valid: true}) {
t.Fatalf("valid license state should trigger automatic license heartbeat")
if licenseHeartbeatEligible(service.LicenseActivationState{Valid: true}) {
t.Fatalf("valid state without stored license key should not trigger automatic license heartbeat")
}
}
func TestLicenseStatusSkipsUnlicensedHeartbeatWithDefaultServer(t *testing.T) {
upstreamCalls := 0
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
upstreamCalls++
w.WriteHeader(http.StatusInternalServerError)
}))
defer upstream.Close()
svc := newLicenseHandlerTestService(t)
svc.Cfg.License.ServerURL = upstream.URL
svc.Cfg.License.HMACSecret = "test-secret"
router := gin.New()
router.GET("/license/status", licenseStatusHandler(svc))
req := httptest.NewRequest(http.MethodGet, "/license/status", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("status = %d body=%s", w.Code, w.Body.String())
}
if upstreamCalls != 0 {
t.Fatalf("unlicensed status should not contact license server, got %d calls", upstreamCalls)
}
}