diff --git a/internal/steward/engine.go b/internal/steward/engine.go index 92b4282..85027ac 100644 --- a/internal/steward/engine.go +++ b/internal/steward/engine.go @@ -7,12 +7,16 @@ import ( ) const safeFailureSummary = "The operation failed locally. No secret-bearing diagnostic was returned through the agent surface." +const stagedFailureSummary = "The operation stopped after a partial state change. Use stage, state_changed, and retry to continue safely." type operationStageError struct { - Stage string - StateChanged string - Retry string - Err error + Stage string + StateChanged string + Retry string + LocalRenderState string + PublicationState string + ClientRefreshState string + Err error } func (e *operationStageError) Error() string { return e.Err.Error() } @@ -146,9 +150,23 @@ func executeReady(ctx context.Context, state *State, request Request) (any, stri data := map[string]any{"summary": safeFailureSummary, "operation": request.Operation} var staged *operationStageError if errors.As(err, &staged) { + data["summary"] = stagedFailureSummary data["stage"] = staged.Stage data["state_changed"] = staged.StateChanged - data["retry"] = staged.Retry + retry := staged.Retry + if request.Operation == "rotate-subscription-token" && staged.StateChanged == "subscription-published-unverified" { + retry = "rotate-subscription-token" + } + data["retry"] = retry + if staged.LocalRenderState != "" { + data["local_render_state"] = staged.LocalRenderState + } + if staged.PublicationState != "" { + data["publication_state"] = staged.PublicationState + } + if staged.ClientRefreshState != "" { + data["client_refresh_state"] = staged.ClientRefreshState + } } return data, code, 1 } diff --git a/internal/steward/migration.go b/internal/steward/migration.go index ec4cbcb..d038000 100644 --- a/internal/steward/migration.go +++ b/internal/steward/migration.go @@ -20,49 +20,59 @@ type migrationState struct { Transactions []migrationTransaction `json:"transactions"` } +type migrationPublicationProgress struct { + TargetID string `json:"target_id"` + Direction string `json:"direction"` + InputFingerprint string `json:"input_fingerprint,omitempty"` + State string `json:"state"` + FailureStage string `json:"failure_stage,omitempty"` +} + type migrationTransaction struct { - ID string `json:"id"` - SourceRoute string `json:"source_route"` - ReplacedServer string `json:"replaced_server"` - ReplacementServer string `json:"replacement_server"` - ReplacementRoute string `json:"replacement_route"` - ReplacementLink *string `json:"replacement_link,omitempty"` - ReplacementServerContext map[string]any `json:"replacement_server_context,omitempty"` - Reason string `json:"reason"` - Phase string `json:"phase"` - Attempt int `json:"attempt"` - CreatedAt string `json:"created_at"` - UpdatedAt string `json:"updated_at"` - AffectedClientTargets []string `json:"affected_client_targets"` - PublicationAttempted []string `json:"publication_attempted"` - CreatedReplacementServer bool `json:"created_replacement_server"` - CreatedReplacementLink bool `json:"created_replacement_link"` - CreatedReplacementRoute bool `json:"created_replacement_route"` - LastFailure string `json:"last_failure,omitempty"` - OldCapacityRetired bool `json:"old_capacity_retired"` - ListenPort int `json:"listen_port"` - PortHopping *PortHopping `json:"port_hopping,omitempty"` - DisplayName string `json:"display_name"` + ID string `json:"id"` + SourceRoute string `json:"source_route"` + ReplacedServer string `json:"replaced_server"` + ReplacementServer string `json:"replacement_server"` + ReplacementRoute string `json:"replacement_route"` + ReplacementLink *string `json:"replacement_link,omitempty"` + ReplacementServerContext map[string]any `json:"replacement_server_context,omitempty"` + Reason string `json:"reason"` + Phase string `json:"phase"` + Attempt int `json:"attempt"` + CreatedAt string `json:"created_at"` + UpdatedAt string `json:"updated_at"` + AffectedClientTargets []string `json:"affected_client_targets"` + PublicationAttempted []string `json:"publication_attempted"` + PublicationProgress []migrationPublicationProgress `json:"publication_progress,omitempty"` + CreatedReplacementServer bool `json:"created_replacement_server"` + CreatedReplacementLink bool `json:"created_replacement_link"` + CreatedReplacementRoute bool `json:"created_replacement_route"` + LastFailure string `json:"last_failure,omitempty"` + OldCapacityRetired bool `json:"old_capacity_retired"` + ListenPort int `json:"listen_port"` + PortHopping *PortHopping `json:"port_hopping,omitempty"` + DisplayName string `json:"display_name"` } type MigrationResult struct { - SchemaVersion int `json:"schema_version"` - MigrationID string `json:"migration_id"` - SourceRoute string `json:"source_route"` - ReplacementServer string `json:"replacement_server"` - ReplacementRoute string `json:"replacement_route"` - ReplacementLink *string `json:"replacement_link,omitempty"` - Phase string `json:"phase"` - Status string `json:"status"` - Summary string `json:"summary"` - AffectedClientTargets []string `json:"affected_client_targets"` - Working map[string]string `json:"working"` - Changed []string `json:"changed"` - LastFailure string `json:"last_failure,omitempty"` - PortHopping *PortHopping `json:"port_hopping,omitempty"` - Next []string `json:"next"` - OldCapacityRetired bool `json:"old_capacity_retired"` - RetirementRequiresAction bool `json:"retirement_requires_explicit_action"` + SchemaVersion int `json:"schema_version"` + MigrationID string `json:"migration_id"` + SourceRoute string `json:"source_route"` + ReplacementServer string `json:"replacement_server"` + ReplacementRoute string `json:"replacement_route"` + ReplacementLink *string `json:"replacement_link,omitempty"` + Phase string `json:"phase"` + Status string `json:"status"` + Summary string `json:"summary"` + AffectedClientTargets []string `json:"affected_client_targets"` + Publication []map[string]string `json:"publication,omitempty"` + Working map[string]string `json:"working"` + Changed []string `json:"changed"` + LastFailure string `json:"last_failure,omitempty"` + PortHopping *PortHopping `json:"port_hopping,omitempty"` + Next []string `json:"next"` + OldCapacityRetired bool `json:"old_capacity_retired"` + RetirementRequiresAction bool `json:"retirement_requires_explicit_action"` } type migrationDependencies struct { @@ -152,9 +162,14 @@ func migrateRouteWith(ctx context.Context, state *State, sourceRoute string, inp return migrationResult(*txn), nil } } - if txn.Phase == "switching" || txn.Phase == "rollback-pending" { + if txn.Phase == "switching" { + txn.Phase = "rollback-pending" + if err := saveMigrationState(state.PrivateDir, store, txn, deps.Now); err != nil { + return MigrationResult{}, err + } + } + if txn.Phase == "rollback-pending" { if err := rollbackMigrationSwitch(state, store, txn, deps); err != nil { - txn.Phase = "rollback-pending" txn.LastFailure = "client-switch-rollback-failed" if saveErr := saveMigrationState(state.PrivateDir, store, txn, deps.Now); saveErr != nil { return MigrationResult{}, errors.Join(err, saveErr) @@ -163,6 +178,7 @@ func migrateRouteWith(ctx context.Context, state *State, sourceRoute string, inp } txn.Phase = "replacement-deployed" txn.PublicationAttempted = []string{} + txn.PublicationProgress = []migrationPublicationProgress{} if err := saveMigrationState(state.PrivateDir, store, txn, deps.Now); err != nil { return MigrationResult{}, err } @@ -229,17 +245,23 @@ func migrateRouteWith(ctx context.Context, state *State, sourceRoute string, inp } txn.Phase = "switching" txn.PublicationAttempted = []string{} + txn.PublicationProgress = []migrationPublicationProgress{} if err := saveMigrationState(state.PrivateDir, store, txn, deps.Now); err != nil { return MigrationResult{}, err } failure, err := switchMigrationClients(state, store, txn, deps) if err != nil { + txn.Phase = "rollback-pending" + txn.LastFailure = failure + if saveErr := saveMigrationState(state.PrivateDir, store, txn, deps.Now); saveErr != nil { + return MigrationResult{}, errors.Join(err, saveErr) + } if rollbackErr := rollbackMigrationSwitch(state, store, txn, deps); rollbackErr != nil { - txn.Phase = "rollback-pending" txn.LastFailure = "client-switch-rollback-failed" } else { txn.Phase = "replacement-deployed" txn.PublicationAttempted = []string{} + txn.PublicationProgress = []migrationPublicationProgress{} txn.LastFailure = failure } if saveErr := saveMigrationState(state.PrivateDir, store, txn, deps.Now); saveErr != nil { @@ -338,7 +360,7 @@ func newMigrationTransaction(state *State, sourceRoute string, input map[string] ReplacedServer: replacedServer, ReplacementServer: replacementServer, ReplacementRoute: replacementRoute, ReplacementLink: replacementLink, ReplacementServerContext: serverContext, Reason: defaultString(stringField(input, "reason"), "planned-replacement"), Phase: "planned", - CreatedAt: timestamp, UpdatedAt: timestamp, AffectedClientTargets: []string{}, PublicationAttempted: []string{}, + CreatedAt: timestamp, UpdatedAt: timestamp, AffectedClientTargets: []string{}, PublicationAttempted: []string{}, PublicationProgress: []migrationPublicationProgress{}, ListenPort: port, PortHopping: portHopping, DisplayName: displayName, }, nil } @@ -444,19 +466,33 @@ func switchMigrationClients(state *State, store *migrationState, txn *migrationT } target := findClientTarget(state.Inventory, targetID) if target != nil && target.Delivery == "subscription" { - txn.PublicationAttempted = sortedUnique(append(txn.PublicationAttempted, targetID)) + fingerprint, err := subscriptionInputFingerprint(state, targetID) + if err != nil { + return "subscription-publication-fingerprint-failed", err + } + setMigrationPublicationProgress(txn, targetID, "forward", fingerprint, "intent", "") if err := saveMigrationState(state.PrivateDir, store, txn, deps.Now); err != nil { return "migration-checkpoint-failed", err } if _, err := deps.Publish(state, targetID, nil); err != nil { + publicationState, failureStage := migrationPublicationFailureState(err) + setMigrationPublicationProgress(txn, targetID, "forward", fingerprint, publicationState, failureStage) + if saveErr := saveMigrationState(state.PrivateDir, store, txn, deps.Now); saveErr != nil { + return "migration-checkpoint-failed", errors.Join(err, saveErr) + } return "subscription-publication-failed", err } + setMigrationPublicationProgress(txn, targetID, "forward", fingerprint, "complete", "") + if err := saveMigrationState(state.PrivateDir, store, txn, deps.Now); err != nil { + return "migration-checkpoint-failed", err + } } } return "", nil } func rollbackMigrationSwitch(state *State, store *migrationState, txn *migrationTransaction, deps migrationDependencies) error { + rollbackTargets := migrationPublicationTargets(txn) if err := setMigrationSelection(state, txn, true); err != nil { return err } @@ -466,16 +502,104 @@ func rollbackMigrationSwitch(state *State, store *migrationState, txn *migration failures = append(failures, err) } } - for _, targetID := range txn.PublicationAttempted { + for _, targetID := range rollbackTargets { + target := findClientTarget(state.Inventory, targetID) + if target == nil || target.Delivery != "subscription" { + continue + } + fingerprint, err := subscriptionInputFingerprint(state, targetID) + if err != nil { + failures = append(failures, err) + continue + } + if progress := migrationPublicationFor(txn, targetID); progress != nil && progress.Direction == "rollback" && progress.State == "complete" && progress.InputFingerprint == fingerprint { + continue + } + setMigrationPublicationProgress(txn, targetID, "rollback", fingerprint, "intent", "") + if err := saveMigrationState(state.PrivateDir, store, txn, deps.Now); err != nil { + failures = append(failures, err) + continue + } if _, err := deps.Publish(state, targetID, nil); err != nil { + publicationState, failureStage := migrationPublicationFailureState(err) + setMigrationPublicationProgress(txn, targetID, "rollback", fingerprint, publicationState, failureStage) + if saveErr := saveMigrationState(state.PrivateDir, store, txn, deps.Now); saveErr != nil { + failures = append(failures, errors.Join(err, saveErr)) + } else { + failures = append(failures, err) + } + continue + } + setMigrationPublicationProgress(txn, targetID, "rollback", fingerprint, "complete", "") + if err := saveMigrationState(state.PrivateDir, store, txn, deps.Now); err != nil { failures = append(failures, err) } } if len(failures) > 0 { return errors.Join(failures...) } - txn.PublicationAttempted = []string{} - return saveMigrationState(state.PrivateDir, store, txn, deps.Now) + return nil +} + +func migrationPublicationFailureState(err error) (string, string) { + var staged *operationStageError + if !errors.As(err, &staged) { + return "intent", "" + } + switch staged.StateChanged { + case "subscription-published-unverified", "new-token-active-at-worker": + return "published-unverified", staged.Stage + case "subscription-published-verified", "subscription-token-rotated": + return "verified", staged.Stage + default: + return "intent", staged.Stage + } +} + +func setMigrationPublicationProgress(txn *migrationTransaction, targetID, direction, fingerprint, state, failureStage string) { + for i := range txn.PublicationProgress { + if txn.PublicationProgress[i].TargetID == targetID { + txn.PublicationProgress[i] = migrationPublicationProgress{TargetID: targetID, Direction: direction, InputFingerprint: fingerprint, State: state, FailureStage: failureStage} + return + } + } + txn.PublicationProgress = append(txn.PublicationProgress, migrationPublicationProgress{TargetID: targetID, Direction: direction, InputFingerprint: fingerprint, State: state, FailureStage: failureStage}) + sort.Slice(txn.PublicationProgress, func(i, j int) bool { return txn.PublicationProgress[i].TargetID < txn.PublicationProgress[j].TargetID }) +} + +func migrationPublicationFor(txn *migrationTransaction, targetID string) *migrationPublicationProgress { + for i := range txn.PublicationProgress { + if txn.PublicationProgress[i].TargetID == targetID { + return &txn.PublicationProgress[i] + } + } + return nil +} + +func migrationPublicationTargets(txn *migrationTransaction) []string { + targets := append([]string(nil), txn.PublicationAttempted...) + for _, progress := range txn.PublicationProgress { + targets = append(targets, progress.TargetID) + } + return sortedUnique(targets) +} + +func validMigrationPublicationProgress(progress migrationPublicationProgress) bool { + if !stableIDPattern.MatchString(progress.TargetID) || (progress.Direction != "forward" && progress.Direction != "rollback") { + return false + } + switch progress.State { + case "intent", "published-unverified", "verified", "complete": + default: + return false + } + if progress.InputFingerprint != "" { + decoded, err := hex.DecodeString(progress.InputFingerprint) + if err != nil || len(decoded) != sha256.Size { + return false + } + } + return true } func setMigrationSelection(state *State, txn *migrationTransaction, restoreOld bool) error { @@ -589,10 +713,18 @@ func migrationResult(txn migrationTransaction) MigrationResult { if txn.Phase == "complete" { changed = append(changed, "affected-client-outputs-switched") } + publication := make([]map[string]string, 0, len(txn.PublicationProgress)) + for _, progress := range txn.PublicationProgress { + item := map[string]string{"target": progress.TargetID, "direction": progress.Direction, "state": progress.State} + if progress.FailureStage != "" { + item["failure_stage"] = progress.FailureStage + } + publication = append(publication, item) + } return MigrationResult{ SchemaVersion: 1, MigrationID: txn.ID, SourceRoute: txn.SourceRoute, ReplacementServer: txn.ReplacementServer, ReplacementRoute: txn.ReplacementRoute, ReplacementLink: txn.ReplacementLink, Phase: txn.Phase, Status: status, - Summary: summary, AffectedClientTargets: append([]string(nil), txn.AffectedClientTargets...), Working: working, + Summary: summary, AffectedClientTargets: append([]string(nil), txn.AffectedClientTargets...), Publication: publication, Working: working, Changed: changed, LastFailure: txn.LastFailure, PortHopping: clonePortHopping(txn.PortHopping), Next: next, OldCapacityRetired: txn.OldCapacityRetired, RetirementRequiresAction: true, } @@ -620,7 +752,16 @@ func readMigrationStateFile(path string) (*migrationState, error) { if txn.ReplacementLink != nil { idsValid = idsValid && stableIDPattern.MatchString(*txn.ReplacementLink) } - if txn.ID == "" || seen[txn.ID] || !idsValid || !validMigrationPhase(txn.Phase) || txn.ListenPort < 1 || txn.ListenPort > 65535 || validatePortHopping(txn.ListenPort, txn.PortHopping) != nil { + progressValid := true + progressTargets := map[string]bool{} + for _, progress := range txn.PublicationProgress { + if progressTargets[progress.TargetID] || !validMigrationPublicationProgress(progress) { + progressValid = false + break + } + progressTargets[progress.TargetID] = true + } + if txn.ID == "" || seen[txn.ID] || !idsValid || !validMigrationPhase(txn.Phase) || !progressValid || txn.ListenPort < 1 || txn.ListenPort > 65535 || validatePortHopping(txn.ListenPort, txn.PortHopping) != nil { return nil, errors.New("private migration state contains an invalid transaction") } seen[txn.ID] = true diff --git a/internal/steward/migration_publication_test.go b/internal/steward/migration_publication_test.go new file mode 100644 index 0000000..a33725d --- /dev/null +++ b/internal/steward/migration_publication_test.go @@ -0,0 +1,178 @@ +package steward + +import ( + "context" + "errors" + "testing" +) + +func TestMigrationPublicationProgressDistinguishesFailureStages(t *testing.T) { + cases := []struct { + name string + publishErr error + wantState string + wantStage string + }{ + {name: "before external mutation", publishErr: errors.New("synthetic pre-publication failure"), wantState: "intent"}, + {name: "after external mutation", publishErr: &operationStageError{Stage: "subscription-verification", StateChanged: "subscription-published-unverified", Retry: "publish-subscription", Err: errors.New("synthetic verification failure")}, wantState: "published-unverified", wantStage: "subscription-verification"}, + {name: "after verified publication", publishErr: &operationStageError{Stage: "client-import-artifact", StateChanged: "subscription-published-verified", Retry: "render-client", Err: errors.New("synthetic client refresh failure")}, wantState: "verified", wantStage: "client-import-artifact"}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + state, source, input := migrationFixture(t, "direct", true) + store, txn := preparedMigrationSwitch(t, state, source, input) + deps := migrationTestDependencies(nil, "healthy") + deps.Publish = func(*State, string, map[string]any) (map[string]any, error) { + return nil, tc.publishErr + } + if _, err := switchMigrationClients(state, store, txn, deps); err == nil { + t.Fatal("publication failure unexpectedly succeeded") + } + progress := migrationPublicationFor(txn, "desktop") + if progress == nil || progress.Direction != "forward" || progress.State != tc.wantState || progress.FailureStage != tc.wantStage || progress.InputFingerprint == "" { + t.Fatalf("publication failure stage was not checkpointed: %#v", progress) + } + }) + } +} + +func TestMigrationPersistsRollbackIntentBeforePublicationRollback(t *testing.T) { + state, source, input := migrationFixture(t, "direct", true) + deps := migrationTestDependencies(nil, "healthy") + publishCalls := 0 + observedRollbackPending := false + deps.Publish = func(state *State, _ string, _ map[string]any) (map[string]any, error) { + publishCalls++ + if publishCalls == 1 { + return nil, &operationStageError{Stage: "client-import-artifact", StateChanged: "subscription-published-verified", Retry: "render-client", Err: errors.New("synthetic forward failure")} + } + store, err := readMigrationState(state.PrivateDir) + if err != nil { + return nil, err + } + txn := findMigration(store, source.ID) + observedRollbackPending = txn != nil && txn.Phase == "rollback-pending" + return nil, errors.New("synthetic rollback interruption") + } + result, err := migrateRouteWith(context.Background(), state, source.ID, input, deps) + if err != nil || result.Phase != "rollback-pending" || result.LastFailure != "client-switch-rollback-failed" { + t.Fatalf("interrupted rollback was not checkpointed: %#v err=%v", result, err) + } + if !observedRollbackPending { + t.Fatal("rollback publication began before rollback-pending was durable") + } +} + +func TestMigrationRollbackResumesAtIncompletePublicationTarget(t *testing.T) { + state, source, input := migrationFixture(t, "direct", true) + if _, err := AddClientTarget(state, map[string]any{"target_id": "tablet", "profile_id": "primary", "renderer": "shadowrocket", "delivery": "nodes"}); err != nil { + t.Fatal(err) + } + if _, err := initializeSubscriptionState(state, "tablet", "synthetic-worker-tablet", "tablet.example.invalid"); err != nil { + t.Fatal(err) + } + store, txn := preparedMigrationSwitch(t, state, source, input) + if err := setMigrationSelection(state, txn, false); err != nil { + t.Fatal(err) + } + for _, targetID := range txn.AffectedClientTargets { + fingerprint, err := subscriptionInputFingerprint(state, targetID) + if err != nil { + t.Fatal(err) + } + setMigrationPublicationProgress(txn, targetID, "forward", fingerprint, "complete", "") + } + txn.Phase = "rollback-pending" + if err := saveMigrationState(state.PrivateDir, store, txn, migrationTestDependencies(nil, "healthy").Now); err != nil { + t.Fatal(err) + } + + deps := migrationTestDependencies(nil, "healthy") + firstCalls := []string{} + deps.Publish = func(_ *State, targetID string, _ map[string]any) (map[string]any, error) { + firstCalls = append(firstCalls, targetID) + if targetID == "tablet" { + return nil, errors.New("synthetic process interruption") + } + return map[string]any{"published": true}, nil + } + if err := rollbackMigrationSwitch(state, store, txn, deps); err == nil { + t.Fatal("rollback interruption unexpectedly succeeded") + } + if len(firstCalls) != 2 || migrationPublicationFor(txn, "desktop").State != "complete" || migrationPublicationFor(txn, "desktop").Direction != "rollback" || migrationPublicationFor(txn, "tablet").State != "intent" { + t.Fatalf("first rollback did not persist per-target progress: calls=%v progress=%#v", firstCalls, txn.PublicationProgress) + } + + reloaded, err := readMigrationState(state.PrivateDir) + if err != nil { + t.Fatal(err) + } + reloadedTxn := findMigration(reloaded, source.ID) + secondCalls := []string{} + retry := migrationTestDependencies(nil, "healthy") + retry.Publish = func(_ *State, targetID string, _ map[string]any) (map[string]any, error) { + secondCalls = append(secondCalls, targetID) + return map[string]any{"published": true}, nil + } + if err := rollbackMigrationSwitch(state, reloaded, reloadedTxn, retry); err != nil { + t.Fatal(err) + } + if len(secondCalls) != 1 || secondCalls[0] != "tablet" { + t.Fatalf("reloaded rollback repeated completed publication work: %v", secondCalls) + } + if !findRoute(state.Inventory, source.ID).Enabled || findRoute(state.Inventory, reloadedTxn.ReplacementRoute).Enabled { + t.Fatal("rollback did not restore the old desired Route") + } +} + +func TestMigrationReadsLegacyPublicationAttemptCheckpoint(t *testing.T) { + state, source, input := migrationFixture(t, "direct", true) + store, txn := preparedMigrationSwitch(t, state, source, input) + txn.PublicationAttempted = []string{"desktop"} + txn.PublicationProgress = nil + txn.Phase = "rollback-pending" + if err := setMigrationSelection(state, txn, false); err != nil { + t.Fatal(err) + } + if err := saveMigrationState(state.PrivateDir, store, txn, migrationTestDependencies(nil, "healthy").Now); err != nil { + t.Fatal(err) + } + + reloaded, err := readMigrationState(state.PrivateDir) + if err != nil { + t.Fatal(err) + } + reloadedTxn := findMigration(reloaded, source.ID) + publishCalls := 0 + deps := migrationTestDependencies(nil, "healthy") + deps.Publish = func(*State, string, map[string]any) (map[string]any, error) { + publishCalls++ + return map[string]any{"published": true}, nil + } + if err := rollbackMigrationSwitch(state, reloaded, reloadedTxn, deps); err != nil { + t.Fatal(err) + } + if publishCalls != 1 || migrationPublicationFor(reloadedTxn, "desktop") == nil || migrationPublicationFor(reloadedTxn, "desktop").State != "complete" { + t.Fatalf("legacy publication attempt did not recover through the new checkpoint model: calls=%d progress=%#v", publishCalls, reloadedTxn.PublicationProgress) + } +} + +func preparedMigrationSwitch(t *testing.T, state *State, source Route, input map[string]any) (*migrationState, *migrationTransaction) { + t.Helper() + deps := migrationTestDependencies(nil, "healthy") + txn, err := newMigrationTransaction(state, source.ID, input, deps.Now) + if err != nil { + t.Fatal(err) + } + store := &migrationState{Schema: 1, Transactions: []migrationTransaction{txn}} + current := &store.Transactions[0] + if err := prepareMigration(state, current); err != nil { + t.Fatal(err) + } + current.AffectedClientTargets = affectedMigrationTargets(state.Inventory, source.ID) + current.Phase = "switching" + if err := saveMigrationState(state.PrivateDir, store, current, deps.Now); err != nil { + t.Fatal(err) + } + return store, current +} diff --git a/internal/steward/observed.go b/internal/steward/observed.go index 4d2b7a8..9cfb2fe 100644 --- a/internal/steward/observed.go +++ b/internal/steward/observed.go @@ -211,6 +211,11 @@ func DriftReport(state *State) (map[string]any, error) { return nil, err } items = append(items, clientItems...) + publicationItems, err := subscriptionPublicationDrift(state) + if err != nil { + return nil, err + } + items = append(items, publicationItems...) errorsCount, warnings := 0, 0 for _, item := range items { if item["severity"] == "error" { diff --git a/internal/steward/subscription.go b/internal/steward/subscription.go index c3ce4b8..12af0ba 100644 --- a/internal/steward/subscription.go +++ b/internal/steward/subscription.go @@ -32,17 +32,39 @@ var ( ) type subscriptionState struct { - Schema int `json:"schema"` - WorkerName string `json:"worker_name"` - Host string `json:"host"` - Token string `json:"token"` - PendingToken *string `json:"pending_token"` - LastPublishedAt *string `json:"last_published_at"` - RotationStartedAt *string `json:"rotation_started_at,omitempty"` - RotatedAt *string `json:"rotated_at,omitempty"` - Reference string `json:"-"` - Path string `json:"-"` - TargetID string `json:"-"` + Schema int `json:"schema"` + WorkerName string `json:"worker_name"` + Host string `json:"host"` + Token string `json:"token"` + PendingToken *string `json:"pending_token"` + LastPublishedAt *string `json:"last_published_at"` + PublishedInputFingerprint string `json:"published_input_fingerprint,omitempty"` + PublishedBodySHA256 string `json:"published_body_sha256,omitempty"` + RotationStartedAt *string `json:"rotation_started_at,omitempty"` + RotatedAt *string `json:"rotated_at,omitempty"` + Reference string `json:"-"` + Path string `json:"-"` + TargetID string `json:"-"` +} + +type subscriptionDependencies struct { + Deploy func(workerName, host, token, format, body string) error + Verify func(endpoint, format, expected string) error + WriteState func(path string, value any) error + Render func(*State, string, bool) (RenderResult, error) + WriteReference func(*State, string, *subscriptionState) (*Artifact, error) + NewToken func() (string, error) +} + +func defaultSubscriptionDependencies() subscriptionDependencies { + return subscriptionDependencies{ + Deploy: deployWorker, + Verify: verifySubscriptionEndpoint, + WriteState: writeJSONAtomic, + Render: RenderClients, + WriteReference: writeSubscriptionReference, + NewToken: newSubscriptionToken, + } } func AssertSubscriptionBodySize(body string) (int, error) { @@ -249,7 +271,36 @@ func writeSubscriptionReference(state *State, targetID string, subscription *sub return &Artifact{ID: targetID + "-subscription", FileName: filepath.Base(path), RelativePath: "/delivery/" + filepath.Base(path)}, nil } +func subscriptionArtifactRetry(target *ClientTarget) string { + if target != nil && target.Renderer == "shadowrocket" { + return "render-client" + } + return "publish-subscription" +} + +func subscriptionEndpoint(subscription *subscriptionState, token string) string { + return "https://" + subscription.Host + "/s/" + token +} + +func ensureSubscriptionPublished(subscription *subscriptionState, target *ClientTarget, token, body string, deps subscriptionDependencies) (bool, error) { + endpoint := subscriptionEndpoint(subscription, token) + if deps.Verify(endpoint, target.Renderer, body) == nil { + return false, nil + } + if err := deps.Deploy(subscription.WorkerName, subscription.Host, token, target.Renderer, body); err != nil { + return false, err + } + if err := deps.Verify(endpoint, target.Renderer, body); err != nil { + return true, &operationStageError{Stage: "subscription-verification", StateChanged: "subscription-published-unverified", Retry: "publish-subscription", PublicationState: "published-unverified", ClientRefreshState: "unknown", Err: err} + } + return true, nil +} + func PublishSubscription(state *State, targetID string, context map[string]any) (map[string]any, error) { + return publishSubscriptionWith(state, targetID, context, defaultSubscriptionDependencies()) +} + +func publishSubscriptionWith(state *State, targetID string, context map[string]any, deps subscriptionDependencies) (map[string]any, error) { subscription, err := readSubscriptionState(state, targetID) if err != nil { if stringField(context, "worker_name") == "" || stringField(context, "host") == "" { @@ -267,24 +318,35 @@ func PublishSubscription(state *State, targetID string, context map[string]any) if err != nil { return nil, err } + fingerprint, err := subscriptionInputFingerprint(state, targetID) + if err != nil { + return nil, err + } target := findClientTarget(state.Inventory, targetID) - if err := deployWorkerAndVerify(subscription.WorkerName, subscription.Host, subscription.Token, target.Renderer, body); err != nil { + deployed, err := ensureSubscriptionPublished(subscription, target, subscription.Token, body, deps) + if err != nil { return nil, err } now := utcNow() subscription.LastPublishedAt = &now - if err := writeJSONAtomic(subscription.Path, subscription); err != nil { - return nil, &operationStageError{Stage: "subscription-state", StateChanged: "subscription-published", Retry: "publish-subscription", Err: err} + subscription.PublishedInputFingerprint = fingerprint + subscription.PublishedBodySHA256 = subscriptionBodySHA256(body) + if err := deps.WriteState(subscription.Path, subscription); err != nil { + return nil, &operationStageError{Stage: "subscription-state", StateChanged: "subscription-published-verified", Retry: "publish-subscription", LocalRenderState: "current", PublicationState: "current", ClientRefreshState: "pending", Err: err} + } + result := map[string]any{ + "client_target": targetID, "worker": subscription.WorkerName, "published": true, "verified": true, + "publication_fingerprint": fingerprint, "remote_mutation_performed": deployed, + "local_render_state": "current", "publication_state": "current", "client_refresh_state": "current", } - result := map[string]any{"client_target": targetID, "worker": subscription.WorkerName, "published": true, "verified": true} if target.Renderer == "shadowrocket" { - if _, err := RenderClients(state, targetID, true); err != nil { - return nil, &operationStageError{Stage: "client-import-artifact", StateChanged: "subscription-published", Retry: "render-client", Err: err} + if _, err := deps.Render(state, targetID, true); err != nil { + return nil, &operationStageError{Stage: "client-import-artifact", StateChanged: "subscription-published-verified", Retry: "render-client", LocalRenderState: "stale-or-missing", PublicationState: "current", ClientRefreshState: "pending", Err: err} } } else { - artifact, err := writeSubscriptionReference(state, targetID, subscription) + artifact, err := deps.WriteReference(state, targetID, subscription) if err != nil { - return nil, &operationStageError{Stage: "client-import-artifact", StateChanged: "subscription-published", Retry: "publish-subscription", Err: err} + return nil, &operationStageError{Stage: "client-import-artifact", StateChanged: "subscription-published-verified", Retry: "publish-subscription", LocalRenderState: "current", PublicationState: "current", ClientRefreshState: "pending", Err: err} } result["import_artifact"] = artifact } @@ -292,6 +354,10 @@ func PublishSubscription(state *State, targetID string, context map[string]any) } func RotateSubscriptionToken(state *State, targetID string) (map[string]any, error) { + return rotateSubscriptionTokenWith(state, targetID, defaultSubscriptionDependencies()) +} + +func rotateSubscriptionTokenWith(state *State, targetID string, deps subscriptionDependencies) (map[string]any, error) { subscription, err := readSubscriptionState(state, targetID) if err != nil { return nil, err @@ -300,14 +366,14 @@ func RotateSubscriptionToken(state *State, targetID string) (map[string]any, err if subscription.PendingToken != nil { proposed = *subscription.PendingToken } else { - proposed, err = newSubscriptionToken() + proposed, err = deps.NewToken() if err != nil { return nil, err } now := utcNow() subscription.PendingToken = &proposed subscription.RotationStartedAt = &now - if err := writeJSONAtomic(subscription.Path, subscription); err != nil { + if err := deps.WriteState(subscription.Path, subscription); err != nil { return nil, err } } @@ -315,13 +381,22 @@ func RotateSubscriptionToken(state *State, targetID string) (map[string]any, err if err != nil { return nil, err } + fingerprint, err := subscriptionInputFingerprint(state, targetID) + if err != nil { + return nil, err + } target := findClientTarget(state.Inventory, targetID) - if err := deployWorkerAndVerify(subscription.WorkerName, subscription.Host, proposed, target.Renderer, body); err != nil { + _, err = ensureSubscriptionPublished(subscription, target, proposed, body, deps) + if err != nil { + if staged := (*operationStageError)(nil); errors.As(err, &staged) { + staged.Retry = "rotate-subscription-token" + staged.StateChanged = "new-token-active-at-worker" + } return nil, err } var fresh subscriptionState if err := readJSON(subscription.Path, &fresh); err != nil { - return nil, err + return nil, &operationStageError{Stage: "subscription-token-state", StateChanged: "new-token-active-at-worker", Retry: "rotate-subscription-token", LocalRenderState: "current", PublicationState: "current", ClientRefreshState: "pending", Err: err} } if fresh.PendingToken == nil || *fresh.PendingToken != proposed { return nil, errors.New("subscription rotation intent changed during publication") @@ -331,20 +406,37 @@ func RotateSubscriptionToken(state *State, targetID string) (map[string]any, err fresh.PendingToken = nil fresh.RotatedAt = &now fresh.LastPublishedAt = &now - if err := writeJSONAtomic(subscription.Path, &fresh); err != nil { - return nil, &operationStageError{Stage: "subscription-token-state", StateChanged: "new-token-active-at-worker", Retry: "rotate-subscription-token", Err: err} + fresh.PublishedInputFingerprint = fingerprint + fresh.PublishedBodySHA256 = subscriptionBodySHA256(body) + if err := deps.WriteState(subscription.Path, &fresh); err != nil { + return nil, &operationStageError{Stage: "subscription-token-state", StateChanged: "new-token-active-at-worker", Retry: "rotate-subscription-token", LocalRenderState: "current", PublicationState: "current", ClientRefreshState: "pending", Err: err} } if target.Renderer == "shadowrocket" { - if _, err := RenderClients(state, targetID, true); err != nil { - return nil, &operationStageError{Stage: "client-import-artifact", StateChanged: "subscription-token-rotated", Retry: "render-client", Err: err} + if _, err := deps.Render(state, targetID, true); err != nil { + return nil, &operationStageError{Stage: "client-import-artifact", StateChanged: "subscription-token-rotated", Retry: subscriptionArtifactRetry(target), LocalRenderState: "stale-or-missing", PublicationState: "current", ClientRefreshState: "pending", Err: err} } - } else if _, err := writeSubscriptionReference(state, targetID, &fresh); err != nil { - return nil, &operationStageError{Stage: "client-import-artifact", StateChanged: "subscription-token-rotated", Retry: "rotate-subscription-token", Err: err} - } - return map[string]any{"client_target": targetID, "token_rotated": true, "published": true, "verified": true, "old_token_revoked_at_worker": true, "unrelated_route_credentials_changed": false, "unrelated_client_credentials_changed": false}, nil + } else if _, err := deps.WriteReference(state, targetID, &fresh); err != nil { + return nil, &operationStageError{Stage: "client-import-artifact", StateChanged: "subscription-token-rotated", Retry: subscriptionArtifactRetry(target), LocalRenderState: "current", PublicationState: "current", ClientRefreshState: "pending", Err: err} + } + return map[string]any{ + "client_target": targetID, "token_rotated": true, "published": true, "verified": true, + "publication_fingerprint": fingerprint, "old_token_revoked_at_worker": true, + "unrelated_route_credentials_changed": false, "unrelated_client_credentials_changed": false, + "local_render_state": "current", "publication_state": "current", "client_refresh_state": "current", + }, nil } func deployWorkerAndVerify(workerName, host, token, format, body string) error { + if err := deployWorker(workerName, host, token, format, body); err != nil { + return err + } + if err := verifySubscriptionEndpoint(subscriptionEndpoint(&subscriptionState{Host: host}, token), format, body); err != nil { + return &operationStageError{Stage: "subscription-verification", StateChanged: "subscription-published-unverified", Retry: "publish-subscription", PublicationState: "published-unverified", ClientRefreshState: "unknown", Err: err} + } + return nil +} + +func deployWorker(workerName, host, token, format, body string) error { if _, err := AssertSubscriptionBodySize(body); err != nil { return err } @@ -418,7 +510,7 @@ func deployWorkerAndVerify(workerName, host, token, format, body string) error { if err := run(npx, append(base, "--secrets-file", secretPath, "--keep-vars", "--minify", "--strict")...); err != nil { return fmt.Errorf("Cloudflare rejected Worker deployment: %w", err) } - return verifySubscriptionEndpoint("https://"+host+"/s/"+token, format, body) + return nil } func verifySubscriptionEndpoint(endpoint, format, expected string) error { @@ -433,11 +525,12 @@ func verifySubscriptionEndpoint(endpoint, format, expected string) error { return errors.New("private subscription endpoint could not be verified after publication") } defer response.Body.Close() - body, err := io.ReadAll(io.LimitReader(response.Body, 8192)) + limit := int64(subscriptionSecretChunkBytes*subscriptionMaxChunks + 1) + body, err := io.ReadAll(io.LimitReader(response.Body, limit)) if err != nil { return err } - if response.StatusCode != http.StatusOK || string(body) != expected { + if response.StatusCode != http.StatusOK || len(body) > subscriptionSecretChunkBytes*subscriptionMaxChunks || string(body) != expected { return errors.New("private subscription endpoint did not return the exact locally generated body") } if !strings.Contains(strings.ToLower(response.Header.Get("Cache-Control")), "no-store") { diff --git a/internal/steward/subscription_fingerprint.go b/internal/steward/subscription_fingerprint.go new file mode 100644 index 0000000..b4962fd --- /dev/null +++ b/internal/steward/subscription_fingerprint.go @@ -0,0 +1,177 @@ +package steward + +import ( + "bytes" + "crypto/sha256" + "encoding/base64" + "encoding/hex" + "encoding/json" + "fmt" + "os" + "path/filepath" + "sort" +) + +func subscriptionInputFingerprint(state *State, targetID string) (string, error) { + target := findClientTarget(state.Inventory, targetID) + if !subscriptionCapableTarget(target) { + return "", fmt.Errorf("ClientTarget %q does not support private subscription delivery", targetID) + } + profile := findProfile(state.Inventory, target.Profile) + if profile == nil { + return "", fmt.Errorf("unknown Profile %q", target.Profile) + } + + hash := sha256.New() + add := func(name string, value any) error { + data, err := json.Marshal(value) + if err != nil { + return err + } + hash.Write([]byte(name)) + hash.Write([]byte{0}) + hash.Write(data) + hash.Write([]byte{0}) + return nil + } + addBytes := func(name string, data []byte) { + hash.Write([]byte(name)) + hash.Write([]byte{0}) + hash.Write(data) + hash.Write([]byte{0}) + } + + if err := add("target", target); err != nil { + return "", err + } + if err := add("profile", profile); err != nil { + return "", err + } + + allRoutes := contains(profile.IncludeRoutes, "*") + for _, route := range state.Inventory.Routes { + if !route.Enabled || (!allRoutes && !contains(profile.IncludeRoutes, route.ID)) { + continue + } + if err := add("route:"+route.ID, route); err != nil { + return "", err + } + path, err := ResolveSecret(route.PayloadSecretRef, state.PrivateDir, nil) + if err != nil { + return "", err + } + data, err := os.ReadFile(path) + if err != nil { + return "", err + } + addBytes("route-payload:"+route.ID, data) + } + + allProviders := contains(profile.IncludeProviders, "*") + for _, provider := range state.Inventory.Providers { + if !provider.Enabled || (!allProviders && !contains(profile.IncludeProviders, provider.ID)) { + continue + } + if err := add("provider:"+provider.ID, provider); err != nil { + return "", err + } + path, err := ResolveSecret(provider.SourceSecretRef, state.PrivateDir, nil) + if err != nil { + return "", err + } + data, err := os.ReadFile(path) + if err != nil { + return "", err + } + addBytes("provider-source:"+provider.ID, data) + } + + return hex.EncodeToString(hash.Sum(nil)), nil +} + +func subscriptionBodySHA256(body string) string { + sum := sha256.Sum256([]byte(body)) + return hex.EncodeToString(sum[:]) +} + +func subscriptionPublicationDrift(state *State) ([]map[string]any, error) { + targets := append([]ClientTarget(nil), state.Inventory.ClientTargets...) + sort.Slice(targets, func(i, j int) bool { return targets[i].ID < targets[j].ID }) + out := []map[string]any{} + for _, target := range targets { + if target.Delivery != "subscription" || !subscriptionCapableTarget(&target) { + continue + } + subscription, err := readSubscriptionState(state, target.ID) + if err != nil { + return nil, err + } + fingerprint, err := subscriptionInputFingerprint(state, target.ID) + if err != nil { + return nil, err + } + + category, severity, observed := "subscription-publication-current", "info", "current" + if subscription.PublishedInputFingerprint == "" || subscription.LastPublishedAt == nil { + category, severity, observed = "subscription-publication-unknown", "warning", "unknown" + } else if subscription.PublishedInputFingerprint != fingerprint { + category, severity, observed = "subscription-publication-stale", "warning", "stale" + } + item := map[string]any{ + "id": "subscription:" + target.ID, + "target": target.ID, + "state_domain": "publication", + "category": category, + "severity": severity, + "desired": "current-with-canonical-state", + "observed": observed, + } + if subscription.LastPublishedAt != nil { + item["observed_at"] = *subscription.LastPublishedAt + } + out = append(out, item) + + refreshCurrent, err := subscriptionClientRefreshCurrent(state, target, subscription) + if err != nil { + return nil, err + } + refreshCategory, refreshSeverity, refreshObserved := "subscription-client-refresh-current", "info", "current" + if !refreshCurrent { + refreshCategory, refreshSeverity, refreshObserved = "subscription-client-refresh-stale", "warning", "stale-or-missing" + } + out = append(out, map[string]any{ + "id": "subscription-client-refresh:" + target.ID, + "target": target.ID, + "state_domain": "client_refresh", + "category": refreshCategory, + "severity": refreshSeverity, + "desired": "current-with-subscription-state", + "observed": refreshObserved, + }) + } + return out, nil +} + +func subscriptionClientRefreshCurrent(state *State, target ClientTarget, subscription *subscriptionState) (bool, error) { + url := subscriptionEndpoint(subscription, subscription.Token) + if target.Renderer == "mihomo" { + data, err := os.ReadFile(filepath.Join(state.Inventory.Delivery.Directory, target.ID+".subscription.txt")) + if os.IsNotExist(err) { + return false, nil + } + if err != nil { + return false, err + } + return string(data) == url+"\n", nil + } + data, err := os.ReadFile(filepath.Join(state.Inventory.Delivery.Directory, target.ID+".html")) + if os.IsNotExist(err) { + return false, nil + } + if err != nil { + return false, err + } + encodedURL := base64.RawURLEncoding.EncodeToString([]byte(url)) + encodedImport := base64.StdEncoding.EncodeToString([]byte("sub://" + encodedURL)) + return bytes.Contains(data, []byte(encodedImport)), nil +} diff --git a/internal/steward/subscription_publication_test.go b/internal/steward/subscription_publication_test.go new file mode 100644 index 0000000..802135c --- /dev/null +++ b/internal/steward/subscription_publication_test.go @@ -0,0 +1,257 @@ +package steward + +import ( + "errors" + "net/http" + "net/http/httptest" + "strings" + "testing" +) + +func TestVerifySubscriptionEndpointAcceptsSupportedLargeBodies(t *testing.T) { + for _, size := range []int{8192, 8193, subscriptionSecretChunkBytes * subscriptionMaxChunks} { + t.Run(strings.Repeat("x", min(size, 32)), func(t *testing.T) { + body := strings.Repeat("x", size) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Cache-Control", "private, no-store") + w.Header().Set("Content-Type", "text/yaml") + _, _ = w.Write([]byte(body)) + })) + defer server.Close() + if err := verifySubscriptionEndpoint(server.URL, "mihomo", body); err != nil { + t.Fatalf("supported body size %d was rejected: %v", size, err) + } + }) + } +} + +func TestVerifySubscriptionEndpointRejectsMismatchedLargeBody(t *testing.T) { + expected := strings.Repeat("x", 8193) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Cache-Control", "no-store") + w.Header().Set("Content-Type", "application/yaml") + _, _ = w.Write([]byte(strings.Repeat("x", 8192) + "y")) + })) + defer server.Close() + if err := verifySubscriptionEndpoint(server.URL, "mihomo", expected); err == nil { + t.Fatal("mismatched body was accepted") + } +} + +func TestSubscriptionPublicationDriftSeparatesPublicationAndClientRefresh(t *testing.T) { + state, route, subscription := subscriptionPublicationFixture(t) + + items, err := subscriptionPublicationDrift(state) + if err != nil { + t.Fatal(err) + } + if driftCategory(items, "publication") != "subscription-publication-unknown" || driftCategory(items, "client_refresh") != "subscription-client-refresh-stale" { + t.Fatalf("initial subscription state was not separated correctly: %#v", items) + } + + fingerprint, err := subscriptionInputFingerprint(state, "desktop") + if err != nil { + t.Fatal(err) + } + body, _, err := ExportSubscriptionBody(state, "desktop") + if err != nil { + t.Fatal(err) + } + now := utcNow() + subscription.LastPublishedAt = &now + subscription.PublishedInputFingerprint = fingerprint + subscription.PublishedBodySHA256 = subscriptionBodySHA256(body) + if err := writeJSONAtomic(subscription.Path, subscription); err != nil { + t.Fatal(err) + } + if _, err := writeSubscriptionReference(state, "desktop", subscription); err != nil { + t.Fatal(err) + } + + items, err = subscriptionPublicationDrift(state) + if err != nil { + t.Fatal(err) + } + if driftCategory(items, "publication") != "subscription-publication-current" || driftCategory(items, "client_refresh") != "subscription-client-refresh-current" { + t.Fatalf("verified publication and client refresh were not current: %#v", items) + } + + findRoute(state.Inventory, route.ID).DisplayName = "changed-route" + items, err = subscriptionPublicationDrift(state) + if err != nil { + t.Fatal(err) + } + if driftCategory(items, "publication") != "subscription-publication-stale" || driftCategory(items, "client_refresh") != "subscription-client-refresh-current" { + t.Fatalf("publication drift was not independent from the client reference: %#v", items) + } +} + +func TestPublishSubscriptionRetrySkipsCompletedRemoteMutation(t *testing.T) { + state, _, _ := subscriptionPublicationFixture(t) + remoteCurrent := false + deployCalls := 0 + stateWrites := 0 + deps := defaultSubscriptionDependencies() + deps.Verify = func(string, string, string) error { + if remoteCurrent { + return nil + } + return errors.New("remote subscription is not current") + } + deps.Deploy = func(string, string, string, string, string) error { + deployCalls++ + remoteCurrent = true + return nil + } + deps.WriteState = func(path string, value any) error { + stateWrites++ + if stateWrites == 1 { + return errors.New("synthetic local state failure") + } + return writeJSONAtomic(path, value) + } + + _, err := publishSubscriptionWith(state, "desktop", nil, deps) + var staged *operationStageError + if !errors.As(err, &staged) || staged.StateChanged != "subscription-published-verified" || staged.PublicationState != "current" || staged.ClientRefreshState != "pending" { + t.Fatalf("verified remote publication was not reported as a partial local failure: %#v err=%v", staged, err) + } + + result, err := publishSubscriptionWith(state, "desktop", nil, deps) + if err != nil { + t.Fatal(err) + } + if deployCalls != 1 || result["remote_mutation_performed"] != false { + t.Fatalf("retry repeated a completed Worker mutation: calls=%d result=%#v", deployCalls, result) + } +} + +func TestRotateSubscriptionTokenRetryReusesPendingToken(t *testing.T) { + state, _, subscription := subscriptionPublicationFixture(t) + proposed := strings.Repeat("A", 43) + remoteToken := subscription.Token + deployCalls := 0 + newTokenCalls := 0 + stateWrites := 0 + deps := defaultSubscriptionDependencies() + deps.NewToken = func() (string, error) { + newTokenCalls++ + return proposed, nil + } + deps.Verify = func(endpoint, _ string, _ string) error { + if strings.HasSuffix(endpoint, "/"+remoteToken) { + return nil + } + return errors.New("token is not active") + } + deps.Deploy = func(_, _, token, _, _ string) error { + deployCalls++ + remoteToken = token + return nil + } + deps.WriteState = func(path string, value any) error { + stateWrites++ + if stateWrites == 2 { + return errors.New("synthetic token commit failure") + } + return writeJSONAtomic(path, value) + } + + _, err := rotateSubscriptionTokenWith(state, "desktop", deps) + var staged *operationStageError + if !errors.As(err, &staged) || staged.StateChanged != "new-token-active-at-worker" { + t.Fatalf("active pending token was not reported after commit failure: %#v err=%v", staged, err) + } + pending, err := readSubscriptionState(state, "desktop") + if err != nil || pending.PendingToken == nil || *pending.PendingToken != proposed { + t.Fatalf("pending token was not preserved for retry: %#v err=%v", pending, err) + } + + result, err := rotateSubscriptionTokenWith(state, "desktop", deps) + if err != nil { + t.Fatal(err) + } + fresh, err := readSubscriptionState(state, "desktop") + if err != nil { + t.Fatal(err) + } + if newTokenCalls != 1 || deployCalls != 1 || fresh.Token != proposed || fresh.PendingToken != nil || result["token_rotated"] != true { + t.Fatalf("rotation retry did not resume the pending token: tokenCalls=%d deployCalls=%d state=%#v result=%#v", newTokenCalls, deployCalls, fresh, result) + } +} + +func TestRotationArtifactFailureRecoversWithoutAnotherRotation(t *testing.T) { + state, _, subscription := subscriptionPublicationFixture(t) + proposed := strings.Repeat("B", 43) + remoteToken := subscription.Token + deployCalls := 0 + newTokenCalls := 0 + referenceWrites := 0 + deps := defaultSubscriptionDependencies() + deps.NewToken = func() (string, error) { + newTokenCalls++ + return proposed, nil + } + deps.Verify = func(endpoint, _ string, _ string) error { + if strings.HasSuffix(endpoint, "/"+remoteToken) { + return nil + } + return errors.New("token is not active") + } + deps.Deploy = func(_, _, token, _, _ string) error { + deployCalls++ + remoteToken = token + return nil + } + deps.WriteReference = func(state *State, targetID string, current *subscriptionState) (*Artifact, error) { + referenceWrites++ + if referenceWrites == 1 { + return nil, errors.New("synthetic reference failure") + } + return writeSubscriptionReference(state, targetID, current) + } + + _, err := rotateSubscriptionTokenWith(state, "desktop", deps) + var staged *operationStageError + if !errors.As(err, &staged) || staged.StateChanged != "subscription-token-rotated" || staged.Retry != "publish-subscription" { + t.Fatalf("artifact failure did not expose the non-rotation recovery action: %#v err=%v", staged, err) + } + committed, err := readSubscriptionState(state, "desktop") + if err != nil || committed.Token != proposed || committed.PendingToken != nil { + t.Fatalf("rotation was not committed before the artifact failure: %#v err=%v", committed, err) + } + + result, err := publishSubscriptionWith(state, "desktop", nil, deps) + if err != nil { + t.Fatal(err) + } + if newTokenCalls != 1 || deployCalls != 1 || referenceWrites != 2 || result["remote_mutation_performed"] != false { + t.Fatalf("artifact recovery repeated credential or remote work: tokenCalls=%d deployCalls=%d referenceWrites=%d result=%#v", newTokenCalls, deployCalls, referenceWrites, result) + } +} + +func subscriptionPublicationFixture(t *testing.T) (*State, Route, *subscriptionState) { + t.Helper() + state, route := healthFixture(t, "direct", false) + if _, err := AddProfile(state, map[string]any{"profile_id": "desktop-profile", "include_routes": []any{route.ID}}); err != nil { + t.Fatal(err) + } + if _, err := AddClientTarget(state, map[string]any{"target_id": "desktop", "profile_id": "desktop-profile", "renderer": "mihomo"}); err != nil { + t.Fatal(err) + } + subscription, err := initializeSubscriptionState(state, "desktop", "synthetic-worker", "subscription.example.invalid") + if err != nil { + t.Fatal(err) + } + return state, route, subscription +} + +func driftCategory(items []map[string]any, domain string) string { + for _, item := range items { + if item["state_domain"] == domain { + value, _ := item["category"].(string) + return value + } + } + return "" +} diff --git a/version.txt b/version.txt index ccbccc3..c043eea 100644 --- a/version.txt +++ b/version.txt @@ -1 +1 @@ -2.2.0 +2.2.1