diff --git a/internal/server/automatic_tasks.go b/internal/server/automatic_tasks.go index 2a7664e..5a51ae0 100644 --- a/internal/server/automatic_tasks.go +++ b/internal/server/automatic_tasks.go @@ -37,6 +37,13 @@ type automaticTaskExecutionError struct { func (value automaticTaskExecutionError) Error() string { return value.err.Error() } func (value automaticTaskExecutionError) Unwrap() error { return value.err } +type automaticTaskProgress func(string) + +type automaticTaskEnvironmentSnapshot struct { + config store.Device + policy store.CardPolicy +} + type automaticTaskScheduler struct { server *Server ctx context.Context @@ -50,6 +57,14 @@ func (s *Server) StartAutomaticTasks(ctx context.Context) { } scheduler := &automaticTaskScheduler{server: s, ctx: ctx, queues: make(map[string]chan store.AutomaticTaskRun)} s.automaticTasks = scheduler + queued, err := s.store.RecoverAutomaticTaskRuns(ctx, time.Now().UTC()) + if err != nil { + s.logger.Warn("recover automatic tasks", "error", err) + } else { + for _, run := range queued { + scheduler.enqueue(run) + } + } go scheduler.run() } @@ -117,9 +132,14 @@ func (scheduler *automaticTaskScheduler) execute(run store.AutomaticTaskRun) { var output string for attempt := 1; attempt <= task.RetryCount+1; attempt++ { run.Attempts = attempt + run.Output = fmt.Sprintf("第 %d 次尝试:正在检查设备和 eSIM Profile", attempt) _ = scheduler.server.store.UpdateAutomaticTaskRun(context.Background(), run) + progress := func(message string) { + run.Output = fmt.Sprintf("第 %d 次尝试:%s", attempt, message) + _ = scheduler.server.store.UpdateAutomaticTaskRun(context.Background(), run) + } operationContext, cancel := context.WithTimeout(scheduler.ctx, automaticTaskMaxRuntime) - output, err = scheduler.server.executeAutomaticTask(operationContext, task) + output, err = scheduler.server.executeAutomaticTask(operationContext, task, progress) cancel() if err == nil { break @@ -154,13 +174,35 @@ func (scheduler *automaticTaskScheduler) execute(run store.AutomaticTaskRun) { } } -func (s *Server) executeAutomaticTask(ctx context.Context, task store.AutomaticTask) (string, error) { - config, entry, physicalID, err := s.ensureAutomaticTaskProfile(ctx, task) +func (s *Server) executeAutomaticTask(ctx context.Context, task store.AutomaticTask, progress automaticTaskProgress) (output string, err error) { + progress("正在检查设备和 eSIM Profile") + config, entry, physicalID, err := s.ensureAutomaticTaskProfile(ctx, task, progress) if err != nil { return "", err } - networkWasEnabled := config.NetworkEnabled - if err := s.prepareAutomaticTaskEnvironment(ctx, &config, entry, physicalID, task); err != nil { + iccid := strings.TrimSpace(task.ProfileICCID) + policy, policyErr := s.store.CardPolicy(ctx, iccid) + if errors.Is(policyErr, store.ErrNotFound) { + policy = defaultCardPolicy(iccid) + } else if policyErr != nil { + return "", fmt.Errorf("read saved card policy: %w", policyErr) + } + snapshot := automaticTaskEnvironmentSnapshot{config: config, policy: policy} + actionCompleted := false + defer func() { + progress("正在恢复该 Profile 原先保存的卡策略") + if restoreErr := s.restoreAutomaticTaskEnvironment(physicalID, snapshot); restoreErr != nil { + if err == nil && actionCompleted { + output = "" + err = automaticTaskExecutionError{err: fmt.Errorf("task completed but card policy restoration failed: %w", restoreErr), retryable: false} + } else if err == nil { + err = fmt.Errorf("restore card policy: %w", restoreErr) + } else { + err = fmt.Errorf("%w; card policy restoration also failed: %v", err, restoreErr) + } + } + }() + if err := s.prepareAutomaticTaskEnvironment(ctx, &config, entry, physicalID, task, progress); err != nil { return "", err } var payload automaticTaskPayload @@ -169,17 +211,22 @@ func (s *Server) executeAutomaticTask(ctx context.Context, task store.AutomaticT } switch task.TaskType { case "sms": - return s.executeAutomaticSMS(ctx, task, payload) + progress("正在发送短信") + output, err = s.executeAutomaticSMS(ctx, task, payload) case "call": - return s.executeAutomaticCall(ctx, task, payload) + progress("正在发起通话") + output, err = s.executeAutomaticCall(ctx, task, payload) case "public_ip": - return s.executeAutomaticPublicIP(ctx, config, physicalID, task.ProfileICCID, networkWasEnabled) + progress("蜂窝数据已连接,正在查询漫游公网 IP") + output, err = s.executeAutomaticPublicIP(ctx, config, task.ProfileICCID) default: return "", fmt.Errorf("unsupported automatic task type %q", task.TaskType) } + actionCompleted = err == nil + return output, err } -func (s *Server) ensureAutomaticTaskProfile(ctx context.Context, task store.AutomaticTask) (store.Device, device.Device, string, error) { +func (s *Server) ensureAutomaticTaskProfile(ctx context.Context, task store.AutomaticTask, progress automaticTaskProgress) (store.Device, device.Device, string, error) { config, err := s.store.Device(ctx, task.DeviceID) if err != nil { return store.Device{}, device.Device{}, "", fmt.Errorf("read device: %w", err) @@ -191,6 +238,7 @@ func (s *Server) ensureAutomaticTaskProfile(ctx context.Context, task store.Auto if strings.EqualFold(strings.TrimSpace(entry.Snapshot.ICCID), strings.TrimSpace(task.ProfileICCID)) { return config, entry, physicalID, nil } + progress("正在切换到任务指定的 eSIM Profile") if _, err := s.devices.SetFlight(ctx, physicalID, true); err != nil { return store.Device{}, device.Device{}, "", fmt.Errorf("enter airplane mode before profile switch: %w", err) } @@ -212,9 +260,10 @@ func (s *Server) ensureAutomaticTaskProfile(ctx context.Context, task store.Auto return config, entry, physicalID, nil } -func (s *Server) prepareAutomaticTaskEnvironment(ctx context.Context, config *store.Device, entry device.Device, physicalID string, task store.AutomaticTask) error { +func (s *Server) prepareAutomaticTaskEnvironment(ctx context.Context, config *store.Device, entry device.Device, physicalID string, task store.AutomaticTask, progress automaticTaskProgress) error { iccid := strings.TrimSpace(task.ProfileICCID) if task.Environment == "vowifi" { + progress("正在准备 VoWiFi 执行环境") if task.TaskType == "public_ip" { return errors.New("public IP tasks cannot run over VoWiFi") } @@ -256,6 +305,7 @@ func (s *Server) prepareAutomaticTaskEnvironment(ctx context.Context, config *st } } } + progress("正在开启蜂窝无线并启用自动选网") config.VoWiFiEnabled = false config.NetworkEnabled = task.TaskType == "public_ip" if err := s.store.UpsertDevice(ctx, *config); err != nil { @@ -289,6 +339,7 @@ func (s *Server) prepareAutomaticTaskEnvironment(ctx context.Context, config *st if _, err := s.devices.ReRegisterOperator(ctx, physicalID); err != nil { return fmt.Errorf("re-register cellular network: %w", err) } + progress("正在搜索并注册蜂窝网络(漫游注册可能需要数分钟)") if err := s.waitAutomaticCellular(ctx, physicalID, task.TaskType == "public_ip"); err != nil { return err } @@ -296,8 +347,8 @@ func (s *Server) prepareAutomaticTaskEnvironment(ctx context.Context, config *st if !s.developerActive(ctx) { return errors.New("roaming public IP tasks require developer mode") } + progress("已注册蜂窝网络,正在建立数据连接") if _, err := s.devices.SetNetwork(ctx, physicalID, s.cardNetworkRequest(ctx, physicalID, *config, policy, true)); err != nil { - s.rollbackAutomaticNetwork(config.ID, physicalID, iccid, *config) return fmt.Errorf("start roaming data: %w", err) } } @@ -421,10 +472,7 @@ func (s *Server) executeAutomaticCall(ctx context.Context, task store.AutomaticT return fmt.Sprintf("已拨打 %s,将在 %d 秒后自动挂断", payload.Phone, payload.DurationSeconds), nil } -func (s *Server) executeAutomaticPublicIP(ctx context.Context, config store.Device, physicalID, iccid string, networkWasEnabled bool) (string, error) { - if !networkWasEnabled { - defer s.rollbackAutomaticNetwork(config.ID, physicalID, iccid, config) - } +func (s *Server) executeAutomaticPublicIP(ctx context.Context, config store.Device, iccid string) (string, error) { if strings.TrimSpace(config.Interface) == "" { return "", errors.New("device has no cellular network interface") } @@ -436,27 +484,70 @@ func (s *Server) executeAutomaticPublicIP(ctx context.Context, config store.Devi return strings.TrimSpace(fmt.Sprintf("公网 IP %s · %s %s", info.IP, info.CountryCode, info.Region)), nil } -func (s *Server) rollbackAutomaticNetwork(deviceID, physicalID, iccid string, config store.Device) { - cleanupContext, cancel := context.WithTimeout(context.Background(), 30*time.Second) +func (s *Server) restoreAutomaticTaskEnvironment(physicalID string, snapshot automaticTaskEnvironmentSnapshot) error { + cleanupContext, cancel := context.WithTimeout(context.Background(), 60*time.Second) defer cancel() - policy, policyErr := s.store.CardPolicy(cleanupContext, iccid) - if policyErr != nil { - policy = store.CardPolicy{ICCID: iccid, APN: config.APN, IPVersion: "IPV4V6"} - } - if _, err := s.devices.SetNetwork(cleanupContext, physicalID, s.cardNetworkRequest(cleanupContext, physicalID, config, policy, false)); err != nil { - s.logger.Warn("stop one-shot automatic roaming data", "device_id", deviceID) - } - config.NetworkEnabled = false - if err := s.store.UpsertDevice(cleanupContext, config); err != nil { - s.logger.Warn("restore automatic roaming data setting", "device_id", deviceID, "error", err) - } - policy.NetworkEnabled = false - policy.VoWiFiEnabled = false - policy.AirplaneEnabled = false - policy.Source = "automatic_task" + config, policy := snapshot.config, snapshot.policy + desiredNetwork := policy.NetworkEnabled && !policy.VoWiFiEnabled && !policy.AirplaneEnabled + config.APN = policy.APN + config.NetworkEnabled = desiredNetwork + config.VoWiFiEnabled = policy.VoWiFiEnabled + var restoreErrors []error if err := s.store.UpsertCardPolicy(cleanupContext, policy); err != nil { - s.logger.Warn("restore automatic roaming card policy", "device_id", deviceID, "error", err) + restoreErrors = append(restoreErrors, fmt.Errorf("persist card policy: %w", err)) } + if err := s.store.UpsertDevice(cleanupContext, config); err != nil { + restoreErrors = append(restoreErrors, fmt.Errorf("persist device policy: %w", err)) + } + + if policy.VoWiFiEnabled { + if _, err := s.devices.SetNetwork(cleanupContext, physicalID, s.cardNetworkRequest(cleanupContext, physicalID, config, policy, false)); err != nil { + restoreErrors = append(restoreErrors, fmt.Errorf("stop cellular data: %w", err)) + } + if _, err := s.devices.SetFlight(cleanupContext, physicalID, true); err != nil { + restoreErrors = append(restoreErrors, fmt.Errorf("restore airplane mode: %w", err)) + } + if s.vowifi == nil { + restoreErrors = append(restoreErrors, errors.New("VoWiFi runtime is unavailable")) + } else if state, stateErr := s.vowifi.State(config.ID); stateErr == nil && state.Enabled { + if _, err := s.vowifi.RequestReconnect(config.ID); err != nil { + restoreErrors = append(restoreErrors, fmt.Errorf("restore VoWiFi: %w", err)) + } + } else if _, err := s.vowifi.RequestEnabled(config.ID, true); err != nil { + restoreErrors = append(restoreErrors, fmt.Errorf("restore VoWiFi: %w", err)) + } + return errors.Join(restoreErrors...) + } + if s.vowifi != nil { + if state, stateErr := s.vowifi.State(config.ID); stateErr == nil && (state.Enabled || state.Active) { + if _, err := s.vowifi.RequestEnabled(config.ID, false); err != nil { + restoreErrors = append(restoreErrors, fmt.Errorf("stop VoWiFi: %w", err)) + } + } + } + if policy.AirplaneEnabled { + if _, err := s.devices.SetNetwork(cleanupContext, physicalID, s.cardNetworkRequest(cleanupContext, physicalID, config, policy, false)); err != nil { + restoreErrors = append(restoreErrors, fmt.Errorf("stop cellular data: %w", err)) + } + if _, err := s.devices.SetFlight(cleanupContext, physicalID, true); err != nil { + restoreErrors = append(restoreErrors, fmt.Errorf("restore airplane mode: %w", err)) + } + return errors.Join(restoreErrors...) + } + if !desiredNetwork { + if _, err := s.devices.SetNetwork(cleanupContext, physicalID, s.cardNetworkRequest(cleanupContext, physicalID, config, policy, false)); err != nil { + restoreErrors = append(restoreErrors, fmt.Errorf("stop cellular data: %w", err)) + } + } + if _, err := s.devices.SetFlight(cleanupContext, physicalID, false); err != nil { + restoreErrors = append(restoreErrors, fmt.Errorf("restore cellular radio: %w", err)) + } + if desiredNetwork { + if _, err := s.devices.SetNetwork(cleanupContext, physicalID, s.cardNetworkRequest(cleanupContext, physicalID, config, policy, true)); err != nil { + restoreErrors = append(restoreErrors, fmt.Errorf("restore cellular data: %w", err)) + } + } + return errors.Join(restoreErrors...) } func (s *Server) cardNetworkRequest( diff --git a/internal/store/automatic_tasks.go b/internal/store/automatic_tasks.go index 391197a..e052e72 100644 --- a/internal/store/automatic_tasks.go +++ b/internal/store/automatic_tasks.go @@ -179,6 +179,43 @@ func (s *Store) UpdateAutomaticTaskRun(ctx context.Context, run AutomaticTaskRun return err } +// RecoverAutomaticTaskRuns reconciles durable run records with the in-memory +// scheduler after a process restart. Running work cannot still be executing, +// while queued work is safe to put back onto the per-device queues. +func (s *Store) RecoverAutomaticTaskRuns(ctx context.Context, now time.Time) ([]AutomaticTaskRun, error) { + const restartError = "service restarted before the automatic task completed" + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return nil, err + } + defer tx.Rollback() + if _, err = tx.ExecContext(ctx, `UPDATE automatic_task_runs SET + status = 'failed', finished_at = ?, error = ?, updated_at = ? + WHERE status = 'running'`, now.Unix(), restartError, now.Unix()); err != nil { + return nil, fmt.Errorf("recover running automatic tasks: %w", err) + } + if _, err = tx.ExecContext(ctx, `UPDATE automatic_tasks SET + last_run_at = ?, last_status = 'failed', last_error = ?, updated_at = ? + WHERE id IN ( + SELECT task_id FROM automatic_task_runs + WHERE status = 'failed' AND error = ? AND finished_at = ? + )`, now.Unix(), restartError, now.Unix(), restartError, now.Unix()); err != nil { + return nil, fmt.Errorf("recover automatic task status: %w", err) + } + rows, err := tx.QueryContext(ctx, automaticTaskRunSelect+` WHERE status = 'queued' ORDER BY id`) + if err != nil { + return nil, fmt.Errorf("recover queued automatic tasks: %w", err) + } + queued, err := scanAutomaticTaskRuns(rows) + if err != nil { + return nil, err + } + if err := tx.Commit(); err != nil { + return nil, err + } + return queued, nil +} + const automaticTaskRunSelect = ` SELECT id, task_id, device_id, scheduled_at, started_at, finished_at, status, attempts, output, error, created_at, updated_at diff --git a/internal/store/automatic_tasks_test.go b/internal/store/automatic_tasks_test.go index ed30b96..7b75500 100644 --- a/internal/store/automatic_tasks_test.go +++ b/internal/store/automatic_tasks_test.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "path/filepath" + "strings" "testing" "time" ) @@ -118,3 +119,64 @@ func TestListAutomaticTaskRunsPaginated(t *testing.T) { t.Fatalf("clamped page: total = %d, runs = %+v", total, all) } } + +func TestRecoverAutomaticTaskRunsFailsRunningAndReturnsQueued(t *testing.T) { + ctx := context.Background() + database := openTestStore(t, filepath.Join(t.TempDir(), "automatic-task-recovery.db")) + mustSaveDevice(t, database, "ec20", "EC20") + task, err := database.SaveAutomaticTask(ctx, AutomaticTask{ + Name: "task", Enabled: true, DeviceID: "ec20", ProfileICCID: "one", + TaskType: "call", Environment: "cellular", IntervalDays: 1, + StartDate: "2026-08-10", RunTime: "12:00", Timezone: "Asia/Shanghai", Payload: []byte(`{"phone":"10086","duration_seconds":10}`), + NextRunAt: time.Now().Add(time.Hour), + }) + if err != nil { + t.Fatal(err) + } + running, err := database.QueueAutomaticTaskNow(ctx, task) + if err != nil { + t.Fatal(err) + } + running.Status = "running" + running.StartedAt = time.Now().UTC().Add(-time.Minute) + running.Attempts = 1 + if err := database.UpdateAutomaticTaskRun(ctx, running); err != nil { + t.Fatal(err) + } + queued, err := database.QueueAutomaticTaskNow(ctx, task) + if err != nil { + t.Fatal(err) + } + + recoveredAt := time.Now().UTC().Truncate(time.Second) + recovered, err := database.RecoverAutomaticTaskRuns(ctx, recoveredAt) + if err != nil { + t.Fatal(err) + } + if len(recovered) != 1 || recovered[0].ID != queued.ID || recovered[0].Status != "queued" { + t.Fatalf("recovered queued runs = %+v", recovered) + } + runs, err := database.ListAutomaticTaskRuns(ctx, 10) + if err != nil { + t.Fatal(err) + } + foundRunning := false + for _, run := range runs { + if run.ID == running.ID { + foundRunning = true + if run.Status != "failed" || run.FinishedAt.IsZero() || !strings.Contains(run.Error, "service restarted") { + t.Fatalf("recovered running run = %+v", run) + } + } + } + if !foundRunning { + t.Fatal("running run was not found after recovery") + } + recoveredTask, err := database.AutomaticTask(ctx, task.ID) + if err != nil { + t.Fatal(err) + } + if recoveredTask.LastStatus != "failed" || !strings.Contains(recoveredTask.LastError, "service restarted") { + t.Fatalf("recovered task status = %+v", recoveredTask) + } +}