Enforce hard limits and hide unavailable task paths

This commit is contained in:
MengMengCode
2026-08-12 17:24:39 +08:00
parent 79ab0573e0
commit d0fd59a2a4
14 changed files with 217 additions and 44 deletions
+46 -5
View File
@@ -83,7 +83,13 @@ func (scheduler *automaticTaskScheduler) run() {
}
func (scheduler *automaticTaskScheduler) claim() {
runs, err := scheduler.server.store.ClaimDueAutomaticTasks(scheduler.ctx, time.Now().UTC(), 50)
var runs []store.AutomaticTaskRun
var err error
if scheduler.server.developerActive(scheduler.ctx) {
runs, err = scheduler.server.store.ClaimDueAutomaticTasks(scheduler.ctx, time.Now().UTC(), 50)
} else {
runs, err = scheduler.server.store.ClaimDueAvailableAutomaticTasks(scheduler.ctx, time.Now().UTC(), 50)
}
if err != nil {
scheduler.server.logger.Warn("claim automatic tasks", "error", err)
return
@@ -127,6 +133,11 @@ func (scheduler *automaticTaskScheduler) execute(run store.AutomaticTaskRun) {
_ = scheduler.server.store.UpdateAutomaticTaskRun(context.Background(), run)
return
}
if err := validateAutomaticTaskAvailability(scheduler.server.developerActive(scheduler.ctx), task.TaskType, task.Environment); err != nil {
run.Status, run.Error, run.FinishedAt = "failed", err.Error(), time.Now().UTC()
_ = scheduler.server.store.UpdateAutomaticTaskRun(context.Background(), run)
return
}
run.Status, run.StartedAt = "running", time.Now().UTC()
_ = scheduler.server.store.UpdateAutomaticTaskRun(context.Background(), run)
var output string
@@ -175,6 +186,9 @@ func (scheduler *automaticTaskScheduler) execute(run store.AutomaticTaskRun) {
}
func (s *Server) executeAutomaticTask(ctx context.Context, task store.AutomaticTask, progress automaticTaskProgress) (output string, err error) {
if err := validateAutomaticTaskAvailability(s.developerActive(ctx), task.TaskType, task.Environment); err != nil {
return "", automaticTaskExecutionError{err: err, retryable: false}
}
progress("正在检查设备和 eSIM Profile")
config, entry, physicalID, err := s.ensureAutomaticTaskProfile(ctx, task, progress)
if err != nil {
@@ -359,9 +373,6 @@ func (s *Server) prepareAutomaticTaskEnvironment(ctx context.Context, config *st
return err
}
if task.TaskType == "public_ip" {
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 {
return fmt.Errorf("start roaming data: %w", err)
@@ -659,6 +670,15 @@ func (s *Server) handleAutomaticTasks(w http.ResponseWriter, r *http.Request) {
s.writeStoreError(w, err)
return
}
if !s.developerActive(r.Context()) {
visible := tasks[:0]
for _, task := range tasks {
if validateAutomaticTaskAvailability(false, task.TaskType, task.Environment) == nil {
visible = append(visible, task)
}
}
tasks = visible
}
writeJSON(w, http.StatusOK, map[string]any{"data": map[string]any{"tasks": tasks}})
case http.MethodPost:
task, err := s.decodeAutomaticTask(r, 0)
@@ -711,7 +731,14 @@ func (s *Server) handleAutomaticTaskRuns(w http.ResponseWriter, r *http.Request)
query := r.URL.Query()
limit, _ := strconv.Atoi(query.Get("limit"))
offset, _ := strconv.Atoi(query.Get("offset"))
runs, total, err := s.store.ListAutomaticTaskRunsPaginated(r.Context(), limit, offset)
var runs []store.AutomaticTaskRun
var total int
var err error
if s.developerActive(r.Context()) {
runs, total, err = s.store.ListAutomaticTaskRunsPaginated(r.Context(), limit, offset)
} else {
runs, total, err = s.store.ListAvailableAutomaticTaskRunsPaginated(r.Context(), limit, offset)
}
if err != nil {
s.writeStoreError(w, err)
return
@@ -741,6 +768,10 @@ func (s *Server) handleAutomaticTaskRunNow(w http.ResponseWriter, r *http.Reques
writeError(w, http.StatusConflict, "wifi_calling_only_device", err.Error())
return
}
if err := validateAutomaticTaskAvailability(s.developerActive(r.Context()), task.TaskType, task.Environment); err != nil {
writeError(w, http.StatusNotFound, "task_unavailable", err.Error())
return
}
run, err := s.store.QueueAutomaticTaskNow(r.Context(), task)
if err != nil {
s.writeStoreError(w, err)
@@ -789,6 +820,9 @@ func (s *Server) decodeAutomaticTask(r *http.Request, id int64) (store.Automatic
if request.TaskType == "public_ip" && request.Environment != "cellular" {
return store.AutomaticTask{}, errors.New("public IP tasks must use cellular direct mode")
}
if err := validateAutomaticTaskAvailability(s.developerActive(r.Context()), request.TaskType, request.Environment); err != nil {
return store.AutomaticTask{}, err
}
if err := validateAutomaticTaskDeviceCapabilities(selectedDevice, request.TaskType, request.Environment); err != nil {
return store.AutomaticTask{}, err
}
@@ -831,6 +865,13 @@ func (s *Server) decodeAutomaticTask(r *http.Request, id int64) (store.Automatic
return task, nil
}
func validateAutomaticTaskAvailability(available bool, taskType, environment string) error {
if !available && (taskType == "public_ip" || environment == "cellular") {
return errors.New("unsupported task type or environment")
}
return nil
}
func validateAutomaticTaskDeviceCapabilities(config store.Device, taskType, environment string) error {
if config.DeviceType != store.DeviceTypeUSBSIMReader {
return nil
+20
View File
@@ -39,6 +39,26 @@ func TestUSBSIMReaderAutomaticTasksRequireVoWiFi(t *testing.T) {
}
}
func TestAutomaticTaskAvailabilityHidesRestrictedPaths(t *testing.T) {
for _, test := range []struct {
available bool
taskType string
environment string
wantError bool
}{
{false, "sms", "vowifi", false},
{false, "call", "vowifi", false},
{false, "sms", "cellular", true},
{false, "public_ip", "cellular", true},
{true, "public_ip", "cellular", false},
} {
err := validateAutomaticTaskAvailability(test.available, test.taskType, test.environment)
if (err != nil) != test.wantError {
t.Fatalf("availability(%v, %q, %q) = %v", test.available, test.taskType, test.environment, err)
}
}
}
func TestAutomaticSMSRetrySafetyPreventsDuplicateSubmission(t *testing.T) {
unsafe := []byte(`{"data":{"parts_attempted":1,"parts_accepted":1,"retry_safe":false}}`)
if automaticSMSRetrySafe(unsafe) {
+3 -3
View File
@@ -39,14 +39,14 @@ func TestDeveloperSettingsUpdatesGlobalSMSLimit(t *testing.T) {
t.Fatal(err)
}
server := &Server{store: database, developerEnabled: true, logger: regionTestLogger(), maxRequestBodyBytes: 4096}
request := httptest.NewRequest(http.MethodPut, "/api/settings/developer", strings.NewReader(`{"sms_hourly_limit":25}`))
request := httptest.NewRequest(http.MethodPut, "/api/settings/developer", strings.NewReader(`{"sms_hourly_limit":17}`))
request.Header.Set("Content-Type", "application/json")
response := httptest.NewRecorder()
server.handleDeveloperSettings(response, request)
if response.Code != http.StatusOK {
t.Fatalf("status = %d, body=%s", response.Code, response.Body.String())
}
if got := developer.SMSHourlyLimit(ctx, database); got != 25 {
t.Fatalf("SMS hourly limit = %d, want 25", got)
if got := developer.SMSHourlyLimit(ctx, database); got != 17 {
t.Fatalf("SMS hourly limit = %d, want 17", got)
}
}