diff --git a/docs/en/reference/admin-api.md b/docs/en/reference/admin-api.md index eca668d..30dd127 100644 --- a/docs/en/reference/admin-api.md +++ b/docs/en/reference/admin-api.md @@ -328,9 +328,12 @@ Pagination is newest-first. When `next_cursor` is present, pass it as `cursor` t | `POST` | `/api/v1/filecoin/readiness/preflight` | Validate pending Filecoin settings. | | `GET` | `/api/v1/observability/providers` | Provider health data. | | `POST` | `/api/v1/observability/providers/refresh` | Refresh provider health. | +| `POST` | `/api/v1/observability/providers/{provider_id}/upload-speed-test` | Start one 32 MiB upload speed test for an available provider. Returns `202 Accepted` with `task_id`, `404 Not Found` for an unknown provider, or `409 Conflict` if a test is already running or the provider is ineligible. | | `GET` | `/api/v1/observability/data-sets` | Local data set health data. | | `POST` | `/api/v1/observability/data-sets/refresh` | Refresh data set health. | +Provider listings include the optional `upload_speed_test` for the latest manual test. A successful result reports `bytes_per_second`, `duration_ms`, `sample_bytes`, and `tested_at`; if the current `service_url` is missing or differs from the tested URL, the result is `stale` instead of a current speed. Tests run only when requested, and the speed is a single sample, not a guarantee for object uploads. Failed tests cannot be retried through the task retry endpoint; start a new test instead. + ## Settings and S3 Users | Method | Path | Purpose | diff --git a/docs/zh/reference/admin-api.md b/docs/zh/reference/admin-api.md index 1fb8e27..31af241 100644 --- a/docs/zh/reference/admin-api.md +++ b/docs/zh/reference/admin-api.md @@ -328,9 +328,12 @@ curl -s "$ADMIN/api/v1/tasks/acknowledge/preview?type=storage_store" | `POST` | `/api/v1/filecoin/readiness/preflight` | 验证待保存的 Filecoin 设置。 | | `GET` | `/api/v1/observability/providers` | 存储提供方健康数据。 | | `POST` | `/api/v1/observability/providers/refresh` | 刷新存储提供方健康状态。 | +| `POST` | `/api/v1/observability/providers/{provider_id}/upload-speed-test` | 对可用的存储提供方发起一次 32 MiB 上传测速。返回 `202 Accepted` 和 `task_id`;存储提供方不存在时返回 `404 Not Found`;已有测速进行中或存储提供方不满足测速条件时返回 `409 Conflict`。 | | `GET` | `/api/v1/observability/data-sets` | 本地数据集健康数据。 | | `POST` | `/api/v1/observability/data-sets/refresh` | 刷新数据集健康状态。 | +存储提供方列表可选返回最近一次手动测速的 `upload_speed_test`。成功结果包含 `bytes_per_second`、`duration_ms`、`sample_bytes` 和 `tested_at`;当前 `service_url` 缺失或与测速时不同,结果显示为 `stale`,不再作为当前速度。测速只在手动发起时运行,结果是单次样本,不保证实际对象上传速度。失败测速不能通过任务重试接口重试,请重新发起测速。 + ## 设置和 S3 用户 | Method | Path | 用途 | diff --git a/internal/admin/api_observability.go b/internal/admin/api_observability.go index 8c9c742..da862f8 100644 --- a/internal/admin/api_observability.go +++ b/internal/admin/api_observability.go @@ -1,12 +1,20 @@ package admin import ( + "context" + "crypto/rand" + "encoding/hex" + "errors" "net/http" "strconv" "strings" "time" + "github.com/strahe/synaps3/internal/db/repository" + "github.com/strahe/synaps3/internal/model" "github.com/strahe/synaps3/internal/observability" + "github.com/strahe/synaps3/internal/providerbenchmark" + taskengine "github.com/strahe/synaps3/internal/task" idtypes "github.com/strahe/synaps3/internal/types" ) @@ -29,7 +37,7 @@ func (s *Server) handleAPIObservabilityProviders(w http.ResponseWriter, r *http. writeJSON(w, http.StatusInternalServerError, map[string]string{"error": "internal"}) return } - writeJSON(w, http.StatusOK, page) + s.writeProviderObservations(w, r, page) } func (s *Server) handleAPIRefreshObservabilityProviders(w http.ResponseWriter, r *http.Request) { @@ -51,7 +59,118 @@ func (s *Server) handleAPIRefreshObservabilityProviders(w http.ResponseWriter, r writeJSON(w, http.StatusInternalServerError, map[string]string{"error": "internal"}) return } - writeJSON(w, http.StatusOK, page) + s.writeProviderObservations(w, r, page) +} + +type providerUploadSpeedView struct { + State string `json:"state"` + SampleBytes int64 `json:"sample_bytes"` + DurationMS *int64 `json:"duration_ms,omitempty"` + BytesPerSecond *int64 `json:"bytes_per_second,omitempty"` + TestedAt *time.Time `json:"tested_at,omitempty"` + FailureCode *string `json:"failure_code,omitempty"` +} + +type providerObservationWithSpeed struct { + observability.ProviderObservation + UploadSpeedTest *providerUploadSpeedView `json:"upload_speed_test,omitempty"` +} + +type providerPageWithSpeed struct { + Items []providerObservationWithSpeed `json:"items"` + Summary observability.Summary `json:"summary"` + SummarySignal observability.SummarySignal `json:"summary_signal"` + Total int `json:"total"` + Limit int `json:"limit"` + Offset int `json:"offset"` +} + +func (s *Server) writeProviderObservations(w http.ResponseWriter, r *http.Request, page observability.ProviderObservationPage) { + if s.repos == nil || s.repos.ProviderUploadSpeed == nil { + writeJSON(w, http.StatusOK, page) + return + } + ids := make([]string, 0, len(page.Items)) + for _, item := range page.Items { + ids = append(ids, item.Facts.ProviderID.String()) + } + tests, err := s.repos.ProviderUploadSpeed.ListByProviderIDs(r.Context(), ids) + if err != nil { + s.logger.Error("api: failed to list provider upload speed tests", "error", err) + writeJSON(w, http.StatusOK, page) + return + } + items := make([]providerObservationWithSpeed, 0, len(page.Items)) + for _, item := range page.Items { + view := providerObservationWithSpeed{ProviderObservation: item} + if row, ok := tests[item.Facts.ProviderID.String()]; ok { + view.UploadSpeedTest = &providerUploadSpeedView{ + State: string(row.State), SampleBytes: row.SampleBytes, + DurationMS: row.DurationMS, BytesPerSecond: row.BytesPerSecond, TestedAt: row.TestedAt, FailureCode: row.FailureCode, + } + if row.State != providerbenchmark.StateTesting && (item.Facts.ServiceURL == nil || providerbenchmark.URLHash(*item.Facts.ServiceURL) != row.ServiceURLHash) { + view.UploadSpeedTest = &providerUploadSpeedView{State: "stale", SampleBytes: row.SampleBytes, TestedAt: row.TestedAt} + } + } + items = append(items, view) + } + writeJSON(w, http.StatusOK, providerPageWithSpeed{ + Items: items, Summary: page.Summary, + SummarySignal: page.SummarySignal, Total: page.Total, Limit: page.Limit, Offset: page.Offset, + }) +} + +func (s *Server) handleAPIProviderUploadSpeedTest(w http.ResponseWriter, r *http.Request) { + if s.observability == nil || s.taskService == nil || s.repos == nil || s.repos.ProviderUploadSpeed == nil { + writeJSON(w, http.StatusServiceUnavailable, map[string]string{"error": "upload speed testing unavailable"}) + return + } + id, err := idtypes.ParseOnChainID("provider_id", r.PathValue("provider_id")) + if err != nil || id.IsZero() { + writeJSON(w, http.StatusBadRequest, map[string]string{"error": "invalid provider ID"}) + return + } + serviceURL, eligible, err := providerbenchmark.CurrentServiceURL(r.Context(), s.observability, id) + if err != nil { + s.logger.Error("api: failed to load provider", "error", err) + writeJSON(w, http.StatusInternalServerError, map[string]string{"error": "internal"}) + return + } + if !eligible { + page, listErr := s.observability.ListProviderObservations(r.Context(), observability.ListOptions{ProviderID: &id, Limit: 1}) + if listErr != nil { + writeJSON(w, http.StatusInternalServerError, map[string]string{"error": "internal"}) + return + } + if len(page.Items) == 0 { + writeJSON(w, http.StatusNotFound, map[string]string{"error": "provider not found"}) + return + } + writeJSON(w, http.StatusConflict, map[string]string{"error": "provider is not available for testing"}) + return + } + var nonce [16]byte + if _, err := rand.Read(nonce[:]); err != nil { + writeJSON(w, http.StatusInternalServerError, map[string]string{"error": "internal"}) + return + } + input := providerbenchmark.Input{ProviderID: id.String(), ServiceURLHash: providerbenchmark.URLHash(serviceURL)} + taskRow, _, err := s.taskService.EnqueueTx(r.Context(), taskengine.EnqueueRequest{ + Type: model.TaskTypeProviderUploadSpeedTest, IdempotencyKey: "provider-upload-speed:" + id.String() + ":" + hex.EncodeToString(nonce[:]), + Input: input, SubjectType: "provider", SubjectKey: id.String(), + }, func(ctx context.Context, repos *repository.Repositories, taskRow *model.Task, _ bool) error { + return repos.ProviderUploadSpeed.Begin(ctx, id.String(), input.ServiceURLHash, taskRow.ID) + }) + if errors.Is(err, repository.ErrConflict) { + writeJSON(w, http.StatusConflict, map[string]string{"error": "upload speed test already running"}) + return + } + if err != nil { + s.logger.Error("api: failed to start provider upload speed test", "error", err) + writeJSON(w, http.StatusInternalServerError, map[string]string{"error": "internal"}) + return + } + writeJSON(w, http.StatusAccepted, map[string]any{"task_id": taskRow.ID, "state": "testing"}) } func (s *Server) handleAPIObservabilityDataSets(w http.ResponseWriter, r *http.Request) { diff --git a/internal/admin/api_observability_test.go b/internal/admin/api_observability_test.go index 098e2eb..b438764 100644 --- a/internal/admin/api_observability_test.go +++ b/internal/admin/api_observability_test.go @@ -3,15 +3,19 @@ package admin import ( "context" "encoding/json" + "errors" "net/http" "net/http/httptest" "strconv" + "sync" "sync/atomic" "testing" "time" "github.com/strahe/synaps3/internal/db/repository" + "github.com/strahe/synaps3/internal/model" "github.com/strahe/synaps3/internal/observability" + "github.com/strahe/synaps3/internal/providerbenchmark" "github.com/strahe/synaps3/internal/testutil" ) @@ -177,6 +181,181 @@ func TestAPIObservabilityProviders(t *testing.T) { } } +type failingUploadSpeedList struct { + repository.ProviderUploadSpeedRepository +} + +func (failingUploadSpeedList) ListByProviderIDs(context.Context, []string) (map[string]providerbenchmark.Result, error) { + return nil, errors.New("speed results unavailable") +} + +func TestAPIObservabilityProvidersKeepsHealthWhenSpeedResultsFail(t *testing.T) { + checkedAt := time.Now().UTC() + service := observability.NewService(observability.ServiceOptions{ + Store: &observabilityAPIStore{ + providers: []observability.ProviderState{{ + ProviderID: onChainID(t, "101"), Status: observability.StatusAvailable, LastCheckedAt: checkedAt, + }}, + providerLastCheckedAt: &checkedAt, + }, + }) + repos := repository.NewRepositories(testutil.NewTestDB(t)) + repos.ProviderUploadSpeed = failingUploadSpeedList{ProviderUploadSpeedRepository: repos.ProviderUploadSpeed} + srv := &Server{repos: repos, observability: service, logger: testLogger()} + rr := httptest.NewRecorder() + srv.handleAPIObservabilityProviders(rr, httptest.NewRequest(http.MethodGet, "/api/v1/observability/providers", nil)) + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want 200: %s", rr.Code, rr.Body.String()) + } + var page struct { + Items []struct { + Facts observability.ProviderFacts `json:"facts"` + UploadSpeedTest json.RawMessage `json:"upload_speed_test"` + } `json:"items"` + } + if err := json.Unmarshal(rr.Body.Bytes(), &page); err != nil { + t.Fatal(err) + } + if len(page.Items) != 1 || page.Items[0].Facts.ProviderID.String() != "101" || page.Items[0].UploadSpeedTest != nil { + t.Fatalf("provider page = %+v", page.Items) + } +} + +func TestAPIProviderUploadSpeedTestConcurrentAdmission(t *testing.T) { + db := testutil.NewTestFileDB(t) + db.SetMaxOpenConns(4) + repos := repository.NewRepositories(db) + checkedAt := time.Now().UTC() + serviceURL := "https://provider.example" + if err := repos.Observability.ReplaceProviderStates(t.Context(), checkedAt, []observability.ProviderState{{ + ProviderID: onChainID(t, "101"), Status: observability.StatusAvailable, + Active: new(true), HasPDP: new(true), ServiceURL: &serviceURL, + LastCheckedAt: checkedAt, ReasonCodes: []observability.ReasonCode{}, Evidence: map[string]any{}, + }}); err != nil { + t.Fatal(err) + } + srv := &Server{ + repos: repos, observability: observability.NewService(observability.ServiceOptions{Store: repos.Observability}), + taskService: newAdminTestTaskService(t, repos), logger: testLogger(), + } + const callers = 8 + start := make(chan struct{}) + statuses := make(chan int, callers) + var workers sync.WaitGroup + for range callers { + workers.Go(func() { + <-start + req := httptest.NewRequest(http.MethodPost, "/api/v1/observability/providers/101/upload-speed-test", nil) + req.SetPathValue("provider_id", "101") + rr := httptest.NewRecorder() + srv.handleAPIProviderUploadSpeedTest(rr, req) + statuses <- rr.Code + }) + } + close(start) + workers.Wait() + close(statuses) + var accepted, conflicts int + for status := range statuses { + switch status { + case http.StatusAccepted: + accepted++ + case http.StatusConflict: + conflicts++ + default: + t.Errorf("concurrent POST status = %d, want 202 or 409", status) + } + } + if accepted != 1 || conflicts != callers-1 { + t.Fatalf("concurrent admission: accepted=%d, conflicts=%d", accepted, conflicts) + } + page, err := repos.Tasks.List(t.Context(), repository.TaskListFilter{Type: model.TaskTypeProviderUploadSpeedTest}) + if err != nil || len(page.Tasks) != 1 { + t.Fatalf("persisted tasks = %+v, err=%v", page.Tasks, err) + } + row, err := repos.ProviderUploadSpeed.Get(t.Context(), "101") + if err != nil || row == nil || row.State != providerbenchmark.StateTesting || row.ActiveTaskID == nil || *row.ActiveTaskID != page.Tasks[0].ID { + t.Fatalf("active test = %+v, err=%v", row, err) + } +} + +func TestAPIProviderUploadSpeedTestAdmitsOneTaskAndListsResult(t *testing.T) { + db := testutil.NewTestDB(t) + repos := repository.NewRepositories(db) + checkedAt := time.Now().UTC() + serviceURL := "https://provider.example" + if err := repos.Observability.ReplaceProviderStates(t.Context(), checkedAt, []observability.ProviderState{{ + ProviderID: onChainID(t, "101"), Status: observability.StatusAvailable, + Active: new(true), HasPDP: new(true), ServiceURL: &serviceURL, + LastCheckedAt: checkedAt, ReasonCodes: []observability.ReasonCode{}, Evidence: map[string]any{}, + }}); err != nil { + t.Fatal(err) + } + srv := &Server{ + repos: repos, observability: observability.NewService(observability.ServiceOptions{Store: repos.Observability}), + taskService: newAdminTestTaskService(t, repos), logger: testLogger(), + } + request := func() *httptest.ResponseRecorder { + req := httptest.NewRequest(http.MethodPost, "/api/v1/observability/providers/101/upload-speed-test", nil) + req.SetPathValue("provider_id", "101") + rr := httptest.NewRecorder() + srv.handleAPIProviderUploadSpeedTest(rr, req) + return rr + } + if rr := request(); rr.Code != http.StatusAccepted { + t.Fatalf("first POST = %d: %s", rr.Code, rr.Body.String()) + } + if rr := request(); rr.Code != http.StatusConflict { + t.Fatalf("duplicate POST = %d: %s", rr.Code, rr.Body.String()) + } + getSpeed := func() providerUploadSpeedView { + t.Helper() + getReq := httptest.NewRequest(http.MethodGet, "/api/v1/observability/providers", nil) + getRR := httptest.NewRecorder() + srv.handleAPIObservabilityProviders(getRR, getReq) + var page struct { + Items []struct { + UploadSpeedTest providerUploadSpeedView `json:"upload_speed_test"` + } `json:"items"` + } + if err := json.Unmarshal(getRR.Body.Bytes(), &page); err != nil { + t.Fatal(err) + } + if len(page.Items) != 1 { + t.Fatalf("listed providers = %+v", page.Items) + } + return page.Items[0].UploadSpeedTest + } + if got := getSpeed(); got.State != string(providerbenchmark.StateTesting) { + t.Fatalf("active speed test = %+v", got) + } + row, err := repos.ProviderUploadSpeed.Get(t.Context(), "101") + if err != nil || row == nil || row.ActiveTaskID == nil { + t.Fatalf("active speed test row = %+v, %v", row, err) + } + if err := repos.ProviderUploadSpeed.Finish(t.Context(), "101", *row.ActiveTaskID, providerbenchmark.StateSucceeded, + 1000, providerbenchmark.SampleBytes, ""); err != nil { + t.Fatal(err) + } + if got := getSpeed(); got.State != string(providerbenchmark.StateSucceeded) || got.BytesPerSecond == nil { + t.Fatalf("successful speed test = %+v", got) + } + newURL := "https://another-provider.example" + if err := repos.Observability.ReplaceProviderStates(t.Context(), time.Now().UTC(), []observability.ProviderState{{ + ProviderID: onChainID(t, "101"), Status: observability.StatusAvailable, + Active: new(true), HasPDP: new(true), ServiceURL: &newURL, + LastCheckedAt: time.Now().UTC(), ReasonCodes: []observability.ReasonCode{}, Evidence: map[string]any{}, + }}); err != nil { + t.Fatal(err) + } + if got := getSpeed(); got.State != "stale" || got.BytesPerSecond != nil { + t.Fatalf("changed-address speed test = %+v", got) + } + if rr := request(); rr.Code != http.StatusAccepted { + t.Fatalf("new-address POST = %d: %s", rr.Code, rr.Body.String()) + } +} + func TestAPIObservabilityRefreshDataSets(t *testing.T) { var calls int32 service := observability.NewService(observability.ServiceOptions{ diff --git a/internal/admin/api_tasks.go b/internal/admin/api_tasks.go index c43a556..330cb85 100644 --- a/internal/admin/api_tasks.go +++ b/internal/admin/api_tasks.go @@ -141,6 +141,7 @@ func validTaskType(taskType model.TaskType) bool { model.TaskTypeStorageDataSetRetire, model.TaskTypeWalletOperation, model.TaskTypeObservabilityRefresh, + model.TaskTypeProviderUploadSpeedTest, model.TaskTypeGC: return true default: @@ -222,6 +223,8 @@ func taskOperationLabel(taskType model.TaskType) string { return "Process wallet request" case model.TaskTypeObservabilityRefresh: return "Refresh storage health" + case model.TaskTypeProviderUploadSpeedTest: + return "Test provider upload speed" case model.TaskTypeGC: return "Remove expired task records" default: diff --git a/internal/admin/api_tasks_test.go b/internal/admin/api_tasks_test.go index 9d72df1..1094b2e 100644 --- a/internal/admin/api_tasks_test.go +++ b/internal/admin/api_tasks_test.go @@ -54,6 +54,7 @@ func newAdminTestTaskService(t *testing.T, repos *repository.Repositories) *task model.TaskTypeStorageDataSetRetire, model.TaskTypeWalletOperation, model.TaskTypeObservabilityRefresh, + model.TaskTypeProviderUploadSpeedTest, model.TaskTypeGC, } { definition := taskengine.Definition{ diff --git a/internal/admin/server.go b/internal/admin/server.go index 8a16ced..c0aee3f 100644 --- a/internal/admin/server.go +++ b/internal/admin/server.go @@ -284,6 +284,7 @@ func (s *Server) Serve(ctx context.Context, listener net.Listener) error { mux.HandleFunc("GET /api/v1/filecoin/readiness", s.handleAPIFilecoinReadiness) mux.HandleFunc("GET /api/v1/observability/providers", s.handleAPIObservabilityProviders) mux.HandleFunc("POST /api/v1/observability/providers/refresh", s.handleAPIRefreshObservabilityProviders) + mux.HandleFunc("POST /api/v1/observability/providers/{provider_id}/upload-speed-test", s.handleAPIProviderUploadSpeedTest) mux.HandleFunc("GET /api/v1/observability/data-sets", s.handleAPIObservabilityDataSets) mux.HandleFunc("POST /api/v1/observability/data-sets/refresh", s.handleAPIRefreshObservabilityDataSets) if s.s3IAM != nil { diff --git a/internal/app/runtime.go b/internal/app/runtime.go index 7e5d2ea..7489144 100644 --- a/internal/app/runtime.go +++ b/internal/app/runtime.go @@ -122,25 +122,26 @@ func NewRuntime(ctx context.Context, opts RuntimeOptions) (_ *Runtime, err error }) registry := taskengine.NewRegistry() handlers, err := worker.NewTaskHandlers(worker.TaskHandlerDependencies{ - Repositories: repos, - Events: events, - Cache: localCache, - CacheGate: cacheGate, - CacheTracker: accessTracker, - Storage: opts.Filecoin.Storage, - Wallet: opts.Filecoin.Wallet, - Receipts: opts.Filecoin.Receipts, - Terminator: opts.Filecoin.Terminator, - Epochs: opts.Filecoin.Epochs, - Observability: observabilityService, - ParkedPieces: pdpStatusChecker, - EvictionPolicy: evictionPolicy, - MaxCacheBytes: maxCacheBytes, - LRUHighPercent: cfg.Cache.LRUHighWatermarkPercent, - LRULowPercent: cfg.Cache.LRULowWatermarkPercent, - DefaultCopies: cfg.Filecoin.DefaultCopies, - MaxRetries: cfg.Worker.Tasks.MaxRetries, - Logger: logger, + Repositories: repos, + Events: events, + Cache: localCache, + CacheGate: cacheGate, + CacheTracker: accessTracker, + Storage: opts.Filecoin.Storage, + Wallet: opts.Filecoin.Wallet, + Receipts: opts.Filecoin.Receipts, + Terminator: opts.Filecoin.Terminator, + Epochs: opts.Filecoin.Epochs, + Observability: observabilityService, + UploadSpeedProbe: synapse.NewPDPBatchUploadProbe(cfg.Filecoin.AllowPrivateNetworks), + ParkedPieces: pdpStatusChecker, + EvictionPolicy: evictionPolicy, + MaxCacheBytes: maxCacheBytes, + LRUHighPercent: cfg.Cache.LRUHighWatermarkPercent, + LRULowPercent: cfg.Cache.LRULowWatermarkPercent, + DefaultCopies: cfg.Filecoin.DefaultCopies, + MaxRetries: cfg.Worker.Tasks.MaxRetries, + Logger: logger, }) if err != nil { return nil, fmt.Errorf("initializing task handlers: %w", err) diff --git a/internal/db/migrations/2026090101_initial_schema.go b/internal/db/migrations/2026090101_initial_schema.go index ebc8c2a..4b8d307 100644 --- a/internal/db/migrations/2026090101_initial_schema.go +++ b/internal/db/migrations/2026090101_initial_schema.go @@ -1057,6 +1057,22 @@ type observabilityProviderState2026090101 struct { Evidence json.RawMessage `bun:"evidence_json,type:jsonb,notnull"` } +type providerUploadSpeedTest2026090101 struct { + bun.BaseModel `bun:"table:provider_upload_speed_tests"` + + ProviderID string `bun:"type:text,pk"` + State string `bun:"type:text,notnull"` + ServiceURLHash string `bun:"type:text,notnull"` + SampleBytes int64 `bun:",notnull"` + DurationMS *int64 + BytesPerSecond *int64 + TestedAt *time.Time + FailureCode *string `bun:"type:text"` + ActiveTaskID *int64 + CreatedAt time.Time `bun:",notnull"` + UpdatedAt time.Time `bun:",notnull"` +} + type observabilityDataSetState2026090101 struct { bun.BaseModel `bun:"table:observability_data_set_states"` @@ -1094,6 +1110,16 @@ func createObservabilitySchema(ctx context.Context, db bun.IDB) error { "CONSTRAINT chk_observability_provider_status CHECK (status IN ('available', 'degraded', 'unavailable', 'unknown'))", }, }, + { + name: "provider_upload_speed_tests", + model: (*providerUploadSpeedTest2026090101)(nil), + constraints: []string{ + "CONSTRAINT chk_provider_upload_speed_tests_state CHECK (state IN ('testing', 'succeeded', 'failed'))", + "CONSTRAINT chk_provider_upload_speed_tests_identity CHECK (provider_id <> '' AND length(service_url_hash) = 64 AND sample_bytes > 0)", + "CONSTRAINT chk_provider_upload_speed_tests_result CHECK ((state = 'testing' AND active_task_id IS NOT NULL AND duration_ms IS NULL AND bytes_per_second IS NULL AND tested_at IS NULL AND failure_code IS NULL) OR (state = 'succeeded' AND active_task_id IS NULL AND duration_ms > 0 AND bytes_per_second > 0 AND tested_at IS NOT NULL AND failure_code IS NULL) OR (state = 'failed' AND active_task_id IS NULL AND duration_ms IS NULL AND bytes_per_second IS NULL AND tested_at IS NOT NULL AND failure_code IS NOT NULL))", + }, + foreignKeys: []string{"(active_task_id) REFERENCES tasks (id) ON UPDATE RESTRICT ON DELETE RESTRICT"}, + }, { name: "observability_data_set_states", model: (*observabilityDataSetState2026090101)(nil), @@ -1115,6 +1141,7 @@ func createObservabilitySchema(ctx context.Context, db bun.IDB) error { } return createInitialIndexes(ctx, db, initialIndexSpec{name: "idx_observability_provider_states_status", table: "observability_provider_states", columns: []string{"status", "last_checked_at"}}, + initialIndexSpec{name: "idx_provider_upload_speed_tests_active_task", table: "provider_upload_speed_tests", columns: []string{"active_task_id"}, where: "active_task_id IS NOT NULL", unique: true}, initialIndexSpec{name: "idx_observability_data_set_states_bucket_status", table: "observability_data_set_states", columns: []string{"bucket_id", "status", "last_checked_at"}}, initialIndexSpec{name: "idx_observability_data_set_states_provider_status", table: "observability_data_set_states", columns: []string{"provider_id", "status", "last_checked_at"}}, ) diff --git a/internal/db/migrations/migrations.go b/internal/db/migrations/migrations.go index 38d92e5..b184846 100644 --- a/internal/db/migrations/migrations.go +++ b/internal/db/migrations/migrations.go @@ -32,6 +32,7 @@ var initialSchemaTableNames = []string{ "observability_collection_states", "observability_data_set_states", "observability_provider_states", + "provider_upload_speed_tests", "s3_accounts", "storage_cleanup_copies", "storage_commit_attempts", @@ -83,6 +84,13 @@ func ValidateTarget(ctx context.Context, db bun.IDB) error { if !statusURLExists || oldJSONExists { return incompatibleDatabaseError() } + benchmarkExists, err := tableExists(ctx, db, "provider_upload_speed_tests") + if err != nil { + return fmt.Errorf("checking provider upload speed tests: %w", err) + } + if !benchmarkExists { + return incompatibleDatabaseError() + } return nil } diff --git a/internal/db/migrations/schema_enum_alignment_test.go b/internal/db/migrations/schema_enum_alignment_test.go index 26241a4..cbf2eef 100644 --- a/internal/db/migrations/schema_enum_alignment_test.go +++ b/internal/db/migrations/schema_enum_alignment_test.go @@ -41,6 +41,7 @@ func TestClosedLifecycleEnumsMatchAppliedChecks(t *testing.T) { {directory: "observability", typeName: "CollectionType", table: "observability_collection_states", constraint: "chk_observability_collection_type"}, {directory: "observability", typeName: "Status", table: "observability_provider_states", constraint: "chk_observability_provider_status"}, {directory: "observability", typeName: "Status", table: "observability_data_set_states", constraint: "chk_observability_data_set_status"}, + {directory: "providerbenchmark", typeName: "State", table: "provider_upload_speed_tests", constraint: "chk_provider_upload_speed_tests_state"}, {directory: "storagecommit", typeName: "AttemptStatus", table: "storage_commit_attempts", constraint: "chk_storage_commit_attempts_status"}, {directory: "storagepull", typeName: "AttemptStatus", table: "storage_pull_attempts", constraint: "chk_storage_pull_attempts_status"}, {directory: "storagereplacement", typeName: "SelectionMode", table: "storage_replacements", constraint: "chk_storage_replacements_selection_mode"}, diff --git a/internal/db/migrations/schema_integrity_test.go b/internal/db/migrations/schema_integrity_test.go index 3047504..a52b335 100644 --- a/internal/db/migrations/schema_integrity_test.go +++ b/internal/db/migrations/schema_integrity_test.go @@ -26,7 +26,7 @@ import ( ) const ( - initialPortableSchemaFingerprint = "657cbdd208ba7ff26431b90df7f03e40233b0475ab25253508b58922fd828c7e" + initialPortableSchemaFingerprint = "4a57bf16ec5d242f43154dd29183049c592641f4b7acee471f4aa269fb2d7b83" ) func TestMigrationRegistryStartsWithUniqueOrderedBaseline(t *testing.T) { @@ -139,7 +139,7 @@ func TestInitialSchemaContractSQLite(t *testing.T) { "multipart_uploads", "multipart_parts", "storage_contents", "storage_data_sets", "storage_copies", "storage_commit_attempts", "storage_replacements", "storage_pull_attempts", "storage_replacement_items", "storage_cleanup_copies", "wallet_operations", "tasks", - "observability_collection_states", "observability_provider_states", "observability_data_set_states", + "observability_collection_states", "observability_provider_states", "observability_data_set_states", "provider_upload_speed_tests", "task_payloads", "storage_data_set_terminations", } { if exists, err := tableExists(t.Context(), db, table); err != nil || !exists { @@ -148,6 +148,7 @@ func TestInitialSchemaContractSQLite(t *testing.T) { } for _, column := range []struct{ table, name string }{ {"tasks", "claim_generation"}, + {"provider_upload_speed_tests", "active_task_id"}, {"task_payloads", "checkpoint_json"}, {"storage_data_set_terminations", "epoch"}, {"storage_copies", "active_task_id"}, @@ -181,6 +182,7 @@ func TestInitialSchemaContractSQLite(t *testing.T) { "idx_storage_data_sets_bucket_provider_active", "idx_storage_replacements_active_bucket_slot", "idx_wallet_operations_recent", + "idx_provider_upload_speed_tests_active_task", } { if exists, err := indexExists(t.Context(), db, index); err != nil || !exists { t.Errorf("index %s exists=%t err=%v", index, exists, err) diff --git a/internal/db/migrations/schema_model_alignment_test.go b/internal/db/migrations/schema_model_alignment_test.go index cc02c9f..fe9f77a 100644 --- a/internal/db/migrations/schema_model_alignment_test.go +++ b/internal/db/migrations/schema_model_alignment_test.go @@ -16,6 +16,7 @@ import ( "github.com/strahe/synaps3/internal/model" "github.com/strahe/synaps3/internal/observability" + "github.com/strahe/synaps3/internal/providerbenchmark" "github.com/strahe/synaps3/internal/storagecommit" "github.com/strahe/synaps3/internal/storagepull" "github.com/strahe/synaps3/internal/storagereplacement" @@ -79,6 +80,7 @@ func runtimePersistentModels() []any { (*model.WalletOperation)(nil), (*observability.CollectionState)(nil), (*observability.ProviderState)(nil), + (*providerbenchmark.Result)(nil), (*observability.DataSetState)(nil), } } diff --git a/internal/db/repository/interfaces.go b/internal/db/repository/interfaces.go index e9eb76d..0953225 100644 --- a/internal/db/repository/interfaces.go +++ b/internal/db/repository/interfaces.go @@ -7,6 +7,7 @@ import ( "github.com/strahe/synaps3/internal/model" "github.com/strahe/synaps3/internal/observability" + "github.com/strahe/synaps3/internal/providerbenchmark" "github.com/strahe/synaps3/internal/storagecommit" "github.com/strahe/synaps3/internal/storagereplacement" "github.com/strahe/synaps3/internal/types" @@ -752,6 +753,14 @@ type ObservabilityRepository interface { GetDataSetStatesByLocalIDs(ctx context.Context, localIDs []int64) (map[int64]observability.DataSetState, error) } +type ProviderUploadSpeedRepository interface { + Begin(context.Context, string, string, int64) error + Finish(context.Context, string, int64, providerbenchmark.State, int64, int64, string) error + FailActiveTask(context.Context, int64, string) error + Get(context.Context, string) (*providerbenchmark.Result, error) + ListByProviderIDs(context.Context, []string) (map[string]providerbenchmark.Result, error) +} + type CreateWalletOperationInput struct { Type model.WalletOperationType ClientRequestID string diff --git a/internal/db/repository/provider_upload_speed_repo.go b/internal/db/repository/provider_upload_speed_repo.go new file mode 100644 index 0000000..988eba2 --- /dev/null +++ b/internal/db/repository/provider_upload_speed_repo.go @@ -0,0 +1,111 @@ +package repository + +import ( + "context" + "database/sql" + "errors" + "time" + + "github.com/strahe/synaps3/internal/providerbenchmark" + "github.com/uptrace/bun" +) + +type BunProviderUploadSpeedRepo struct{ db bun.IDB } + +func (r *BunProviderUploadSpeedRepo) Begin(ctx context.Context, providerID, hash string, taskID int64) error { + now := time.Now().UTC() + row := &providerbenchmark.Result{ + ProviderID: providerID, State: providerbenchmark.StateTesting, ServiceURLHash: hash, + SampleBytes: providerbenchmark.SampleBytes, ActiveTaskID: &taskID, CreatedAt: now, UpdatedAt: now, + } + result, err := r.db.NewInsert().Model(row).On("CONFLICT (provider_id) DO NOTHING").Exec(ctx) + if err != nil { + return err + } + if count, _ := result.RowsAffected(); count == 1 { + return nil + } + result, err = r.db.NewUpdate().Model((*providerbenchmark.Result)(nil)). + Set("state = ?", providerbenchmark.StateTesting). + Set("service_url_hash = ?", hash). + Set("sample_bytes = ?", providerbenchmark.SampleBytes). + Set("duration_ms = NULL").Set("bytes_per_second = NULL").Set("tested_at = NULL").Set("failure_code = NULL"). + Set("active_task_id = ?", taskID).Set("updated_at = ?", now). + Where("provider_id = ? AND state <> ?", providerID, providerbenchmark.StateTesting).Exec(ctx) + if err != nil { + return err + } + if count, _ := result.RowsAffected(); count != 1 { + return ErrConflict + } + return nil +} + +func (r *BunProviderUploadSpeedRepo) Finish(ctx context.Context, providerID string, taskID int64, state providerbenchmark.State, durationMS, bytesPerSecond int64, failureCode string) error { + if state != providerbenchmark.StateSucceeded && state != providerbenchmark.StateFailed { + return ErrInvalidInput + } + now := time.Now().UTC() + query := r.db.NewUpdate().Model((*providerbenchmark.Result)(nil)). + Set("state = ?", state).Set("active_task_id = NULL").Set("tested_at = ?", now).Set("updated_at = ?", now). + Where("provider_id = ? AND active_task_id = ? AND state = ?", providerID, taskID, providerbenchmark.StateTesting) + if state == providerbenchmark.StateSucceeded { + if durationMS < 1 || bytesPerSecond < 1 { + return ErrInvalidInput + } + query = query.Set("duration_ms = ?", durationMS).Set("bytes_per_second = ?", bytesPerSecond).Set("failure_code = NULL") + } else { + if failureCode == "" { + return ErrInvalidInput + } + query = query.Set("duration_ms = NULL").Set("bytes_per_second = NULL").Set("failure_code = ?", failureCode) + } + result, err := query.Exec(ctx) + if err != nil { + return err + } + if count, _ := result.RowsAffected(); count != 1 { + return ErrConflict + } + return nil +} + +func (r *BunProviderUploadSpeedRepo) FailActiveTask(ctx context.Context, taskID int64, failureCode string) error { + if taskID < 1 || failureCode == "" { + return ErrInvalidInput + } + now := time.Now().UTC() + _, err := r.db.NewUpdate().Model((*providerbenchmark.Result)(nil)). + Set("state = ?", providerbenchmark.StateFailed). + Set("active_task_id = NULL").Set("duration_ms = NULL").Set("bytes_per_second = NULL"). + Set("failure_code = ?", failureCode).Set("tested_at = ?", now).Set("updated_at = ?", now). + Where("active_task_id = ? AND state = ?", taskID, providerbenchmark.StateTesting).Exec(ctx) + return err +} + +func (r *BunProviderUploadSpeedRepo) Get(ctx context.Context, providerID string) (*providerbenchmark.Result, error) { + var row providerbenchmark.Result + err := r.db.NewSelect().Model(&row).Where("provider_id = ?", providerID).Scan(ctx) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + if err != nil { + return nil, err + } + return &row, nil +} + +func (r *BunProviderUploadSpeedRepo) ListByProviderIDs(ctx context.Context, ids []string) (map[string]providerbenchmark.Result, error) { + out := make(map[string]providerbenchmark.Result, len(ids)) + if len(ids) == 0 { + return out, nil + } + var rows []providerbenchmark.Result + if err := r.db.NewSelect().Model(&rows).Where("provider_id IN (?)", bun.List(ids)).Scan(ctx); err != nil { + return nil, err + } + for _, row := range rows { + out[row.ProviderID] = row + } + return out, nil +} diff --git a/internal/db/repository/provider_upload_speed_repo_test.go b/internal/db/repository/provider_upload_speed_repo_test.go new file mode 100644 index 0000000..3a4f36a --- /dev/null +++ b/internal/db/repository/provider_upload_speed_repo_test.go @@ -0,0 +1,98 @@ +package repository_test + +import ( + "errors" + "testing" + "time" + + "github.com/strahe/synaps3/internal/db/repository" + "github.com/strahe/synaps3/internal/model" + "github.com/strahe/synaps3/internal/observability" + "github.com/strahe/synaps3/internal/providerbenchmark" +) + +func TestProviderUploadSpeedKeepsOnlyLatestResultAcrossHealthRefresh(t *testing.T) { + db := testDB(t) + repos := repository.NewRepositories(db) + newTask := func(key string) int64 { + t.Helper() + row, created, err := repos.Tasks.Enqueue(t.Context(), &model.Task{ + Type: model.TaskTypeProviderUploadSpeedTest, IdempotencyKey: key, InputVersion: 1, + Input: []byte(`{}`), InputHash: key, Status: model.TaskStatusPending, + ResumeMode: model.TaskResumeModeExecute, AvailableAt: time.Now(), + }) + if err != nil || !created { + t.Fatalf("enqueue task: %v, created=%v", err, created) + } + return row.ID + } + first := newTask("first-speed-test") + hash := providerbenchmark.URLHash("https://provider.example") + if err := repos.ProviderUploadSpeed.Begin(t.Context(), "101", hash, first); err != nil { + t.Fatal(err) + } + second := newTask("second-speed-test") + if err := repos.ProviderUploadSpeed.Begin(t.Context(), "101", hash, second); !errors.Is(err, repository.ErrConflict) { + t.Fatalf("concurrent Begin = %v, want conflict", err) + } + if err := repos.ProviderUploadSpeed.Finish(t.Context(), "101", first, providerbenchmark.StateSucceeded, 1000, providerbenchmark.SampleBytes, ""); err != nil { + t.Fatal(err) + } + if err := repos.Observability.ReplaceProviderStates(t.Context(), time.Now(), []observability.ProviderState{{ + ProviderID: onChainID(t, "101"), Status: observability.StatusAvailable, + ReasonCodes: []observability.ReasonCode{}, Evidence: map[string]any{}, + }}); err != nil { + t.Fatal(err) + } + row, err := repos.ProviderUploadSpeed.Get(t.Context(), "101") + if err != nil || row == nil || row.State != providerbenchmark.StateSucceeded || row.BytesPerSecond == nil { + t.Fatalf("result after refresh = %#v, %v", row, err) + } + if err := repos.ProviderUploadSpeed.Begin(t.Context(), "101", hash, second); err != nil { + t.Fatal(err) + } + if err := repos.ProviderUploadSpeed.FailActiveTask(t.Context(), first, "handler_panic"); err != nil { + t.Fatal(err) + } + if active, err := repos.ProviderUploadSpeed.Get(t.Context(), "101"); err != nil || active == nil || active.State != providerbenchmark.StateTesting || active.ActiveTaskID == nil || *active.ActiveTaskID != second { + t.Fatalf("new active test after stale task failure = %#v, %v", active, err) + } + if err := repos.ProviderUploadSpeed.Finish(t.Context(), "101", first, providerbenchmark.StateSucceeded, 1000, providerbenchmark.SampleBytes, ""); !errors.Is(err, repository.ErrConflict) { + t.Fatalf("stale task Finish = %v, want conflict", err) + } + if err := repos.ProviderUploadSpeed.Finish(t.Context(), "101", second, providerbenchmark.StateFailed, 0, 0, "timeout"); err != nil { + t.Fatal(err) + } + row, err = repos.ProviderUploadSpeed.Get(t.Context(), "101") + if err != nil || row == nil || row.State != providerbenchmark.StateFailed || row.BytesPerSecond != nil || row.ActiveTaskID != nil || row.FailureCode == nil || *row.FailureCode != "timeout" { + t.Fatalf("latest result = %#v, %v", row, err) + } +} + +func TestProviderUploadSpeedReleasesTaskForRetentionCleanup(t *testing.T) { + db := testDB(t) + repos := repository.NewRepositories(db) + claimed := enqueueAndClaimTask(t, repos, "speed-task-gc", time.Minute) + if err := repos.ProviderUploadSpeed.Begin(t.Context(), "101", providerbenchmark.URLHash("https://provider.example"), claimed.ID); err != nil { + t.Fatal(err) + } + expired := time.Now().Add(-time.Minute) + if err := repos.Tasks.Settle(t.Context(), claimed.ID, claimed.ClaimGeneration, repository.TaskTransition{ + Status: model.TaskStatusCompleted, ResumeMode: model.TaskResumeModeRecover, RetentionUntil: &expired, + }); err != nil { + t.Fatal(err) + } + if deleted, err := repos.Tasks.DeleteRetained(t.Context(), time.Now(), 10); err != nil || deleted != 0 { + t.Fatalf("cleanup while speed test references task = %d, %v", deleted, err) + } + if err := repos.ProviderUploadSpeed.Finish(t.Context(), "101", claimed.ID, providerbenchmark.StateSucceeded, 1000, providerbenchmark.SampleBytes, ""); err != nil { + t.Fatal(err) + } + if deleted, err := repos.Tasks.DeleteRetained(t.Context(), time.Now(), 10); err != nil || deleted != 1 { + t.Fatalf("cleanup after speed test settles = %d, %v", deleted, err) + } + row, err := repos.ProviderUploadSpeed.Get(t.Context(), "101") + if err != nil || row == nil || row.State != providerbenchmark.StateSucceeded { + t.Fatalf("result after task cleanup = %#v, %v", row, err) + } +} diff --git a/internal/db/repository/repos.go b/internal/db/repository/repos.go index 91a1eef..18ecee6 100644 --- a/internal/db/repository/repos.go +++ b/internal/db/repository/repos.go @@ -12,17 +12,18 @@ import ( // dependency for service/backend layers. The WithTx helper executes a callback // inside a database transaction with a clone of Repositories backed by the tx. type Repositories struct { - Buckets BucketRepository - S3Accounts S3AccountRepository - Objects ObjectRepository - Contents StorageContentRepository - Replacements StorageReplacementRepository - StorageCleanup StorageCleanupRepository - Tasks TaskRepository - CacheEvictions CacheEvictionRepository - Multiparts MultipartUploadRepository - WalletOperations WalletOperationRepository - Observability ObservabilityRepository + Buckets BucketRepository + S3Accounts S3AccountRepository + Objects ObjectRepository + Contents StorageContentRepository + Replacements StorageReplacementRepository + StorageCleanup StorageCleanupRepository + Tasks TaskRepository + CacheEvictions CacheEvictionRepository + Multiparts MultipartUploadRepository + WalletOperations WalletOperationRepository + Observability ObservabilityRepository + ProviderUploadSpeed ProviderUploadSpeedRepository db bun.IDB } @@ -30,18 +31,19 @@ type Repositories struct { // NewRepositories constructs a Repositories with concrete Bun-backed implementations. func NewRepositories(db bun.IDB) *Repositories { return &Repositories{ - Buckets: &BunBucketRepo{db: db}, - S3Accounts: &BunS3AccountRepo{db: db}, - Objects: &BunObjectRepo{db: db}, - Contents: &BunStorageContentRepo{db: db}, - Replacements: &BunStorageReplacementRepo{db: db}, - StorageCleanup: &BunStorageCleanupRepo{db: db}, - Tasks: &BunTaskRepo{db: db}, - CacheEvictions: &BunCacheEvictionRepo{db: db}, - Multiparts: &BunMultipartRepo{db: db}, - WalletOperations: &BunWalletOperationRepo{db: db}, - Observability: &BunObservabilityRepo{db: db}, - db: db, + Buckets: &BunBucketRepo{db: db}, + S3Accounts: &BunS3AccountRepo{db: db}, + Objects: &BunObjectRepo{db: db}, + Contents: &BunStorageContentRepo{db: db}, + Replacements: &BunStorageReplacementRepo{db: db}, + StorageCleanup: &BunStorageCleanupRepo{db: db}, + Tasks: &BunTaskRepo{db: db}, + CacheEvictions: &BunCacheEvictionRepo{db: db}, + Multiparts: &BunMultipartRepo{db: db}, + WalletOperations: &BunWalletOperationRepo{db: db}, + Observability: &BunObservabilityRepo{db: db}, + ProviderUploadSpeed: &BunProviderUploadSpeedRepo{db: db}, + db: db, } } diff --git a/internal/db/repository/task_repo.go b/internal/db/repository/task_repo.go index 9f4f06d..71e5814 100644 --- a/internal/db/repository/task_repo.go +++ b/internal/db/repository/task_repo.go @@ -637,6 +637,7 @@ func (r *BunTaskRepo) DeleteRetained(ctx context.Context, now time.Time, limit i Where(`NOT EXISTS (SELECT 1 FROM storage_copies WHERE active_task_id = task.id)`). Where(`NOT EXISTS (SELECT 1 FROM storage_data_sets WHERE ensure_task_id = task.id OR retirement_task_id = task.id)`). Where(`NOT EXISTS (SELECT 1 FROM wallet_operations WHERE task_id = task.id)`). + Where(`NOT EXISTS (SELECT 1 FROM provider_upload_speed_tests WHERE active_task_id = task.id)`). Where(`NOT EXISTS (SELECT 1 FROM storage_replacements WHERE task_id = task.id)`). OrderExpr("retention_until, id"). Limit(limit). diff --git a/internal/model/task.go b/internal/model/task.go index 5f7923e..00de9e6 100644 --- a/internal/model/task.go +++ b/internal/model/task.go @@ -29,6 +29,7 @@ const ( TaskTypeStorageDataSetRetire TaskType = "storage_dataset_retire" TaskTypeWalletOperation TaskType = "wallet_operation" TaskTypeObservabilityRefresh TaskType = "observability_refresh" + TaskTypeProviderUploadSpeedTest TaskType = "provider_upload_speed_test" TaskTypeGC TaskType = "task_gc" ) diff --git a/internal/providerbenchmark/model.go b/internal/providerbenchmark/model.go new file mode 100644 index 0000000..38cf79e --- /dev/null +++ b/internal/providerbenchmark/model.go @@ -0,0 +1,71 @@ +package providerbenchmark + +import ( + "context" + "crypto/sha256" + "encoding/hex" + "time" + + "github.com/strahe/synaps3/internal/observability" + "github.com/strahe/synaps3/internal/types" + "github.com/uptrace/bun" +) + +const SampleBytes int64 = 32 << 20 + +type State string + +const ( + StateTesting State = "testing" + StateSucceeded State = "succeeded" + StateFailed State = "failed" +) + +type Result struct { + bun.BaseModel `bun:"table:provider_upload_speed_tests"` + + ProviderID string `bun:"provider_id,pk,type:text" json:"-"` + State State `bun:"state,type:text,notnull" json:"state"` + ServiceURLHash string `bun:"service_url_hash,type:text,notnull" json:"-"` + SampleBytes int64 `bun:"sample_bytes,notnull" json:"sample_bytes"` + DurationMS *int64 `bun:"duration_ms" json:"duration_ms,omitempty"` + BytesPerSecond *int64 `bun:"bytes_per_second" json:"bytes_per_second,omitempty"` + TestedAt *time.Time `bun:"tested_at" json:"tested_at,omitempty"` + FailureCode *string `bun:"failure_code,type:text" json:"failure_code,omitempty"` + ActiveTaskID *int64 `bun:"active_task_id" json:"-"` + CreatedAt time.Time `bun:"created_at,notnull" json:"-"` + UpdatedAt time.Time `bun:"updated_at,notnull" json:"-"` +} + +type Input struct { + ProviderID string `json:"provider_id"` + ServiceURLHash string `json:"service_url_hash"` +} + +type Checkpoint struct { + Attempted bool `json:"attempted"` + DurationMS int64 `json:"duration_ms,omitempty"` + BytesPerSecond int64 `json:"bytes_per_second,omitempty"` +} + +func URLHash(value string) string { + sum := sha256.Sum256([]byte(value)) + return hex.EncodeToString(sum[:]) +} + +func CurrentServiceURL(ctx context.Context, service *observability.Service, providerID types.OnChainID) (string, bool, error) { + page, err := service.ListProviderObservations(ctx, observability.ListOptions{ProviderID: &providerID, Limit: 1}) + if err != nil { + return "", false, err + } + if len(page.Items) == 0 { + return "", false, nil + } + item := page.Items[0] + if item.Signal.Status != observability.StatusAvailable || item.Signal.Freshness.Stale || + item.Facts.Active == nil || !*item.Facts.Active || item.Facts.HasPDP == nil || !*item.Facts.HasPDP || + item.Facts.ServiceURL == nil || *item.Facts.ServiceURL == "" { + return "", false, nil + } + return *item.Facts.ServiceURL, true, nil +} diff --git a/internal/synapse/pdp_upload_probe.go b/internal/synapse/pdp_upload_probe.go new file mode 100644 index 0000000..df08096 --- /dev/null +++ b/internal/synapse/pdp_upload_probe.go @@ -0,0 +1,112 @@ +package synapse + +import ( + "bytes" + "context" + "crypto/rand" + "errors" + "fmt" + "io" + "net/http" + "net/url" + "regexp" + "strings" + "time" + + "github.com/strahe/synaps3/internal/providerbenchmark" +) + +const uploadProbeTimeout = 180 * time.Second + +var uploadSessionPath = regexp.MustCompile(`(?:^|/)pdp/piece/uploads/([a-fA-F0-9]{8}-[a-fA-F0-9]{4}-[a-fA-F0-9]{4}-[a-fA-F0-9]{4}-[a-fA-F0-9]{12})$`) + +type PDPBatchUploadProbe struct { + client *http.Client + timeout time.Duration +} + +func NewPDPBatchUploadProbe(allowPrivate bool) *PDPBatchUploadProbe { + return &PDPBatchUploadProbe{client: newPDPStatusHTTPClient(0, allowPrivate), timeout: uploadProbeTimeout} +} + +// Probe measures only the PUT phase; the upload session is intentionally not finalized. +func (p *PDPBatchUploadProbe) Probe(ctx context.Context, serviceURL string) (time.Duration, error) { + if err := ctx.Err(); err != nil { + return 0, err + } + base, err := parsePDPHTTPURL(serviceURL, "provider service URL") + if err != nil || base.User != nil || base.RawQuery != "" || base.Fragment != "" { + return 0, errors.New("invalid provider service URL") + } + base = cloneUploadBaseURL(base) + createURL := base.ResolveReference(&url.URL{Path: "pdp/piece/uploads"}) + payload := make([]byte, providerbenchmark.SampleBytes) + if _, err := rand.Read(payload); err != nil { + return 0, fmt.Errorf("preparing upload sample: %w", err) + } + ctx, cancel := context.WithTimeout(ctx, p.timeout) + defer cancel() + createReq, err := http.NewRequestWithContext(ctx, http.MethodPost, createURL.String(), nil) + if err != nil { + return 0, err + } + createResp, err := p.client.Do(createReq) + if err != nil { + return 0, err + } + defer func() { _ = createResp.Body.Close() }() + _, _ = io.CopyN(io.Discard, createResp.Body, 1024) + if createResp.StatusCode != http.StatusCreated { + return 0, fmt.Errorf("upload session creation returned HTTP %d", createResp.StatusCode) + } + uuid, err := checkedUploadSessionID(createResp.Header.Get("Location"), createURL) + if err != nil { + return 0, err + } + putURL := base.ResolveReference(&url.URL{Path: "pdp/piece/uploads/" + uuid}) + // An opaque ReadCloser prevents net/http from constructing a replayable body. + putReq, err := http.NewRequestWithContext(ctx, http.MethodPut, putURL.String(), io.NopCloser(bytes.NewReader(payload))) + if err != nil { + return 0, err + } + putReq.ContentLength = providerbenchmark.SampleBytes + putReq.Header.Set("Content-Type", "application/octet-stream") + started := time.Now() + putResp, err := p.client.Do(putReq) + if err != nil { + return 0, err + } + duration := time.Since(started) + defer func() { _ = putResp.Body.Close() }() + _, _ = io.CopyN(io.Discard, putResp.Body, 1024) + if putResp.StatusCode != http.StatusNoContent { + return 0, fmt.Errorf("upload sample returned HTTP %d", putResp.StatusCode) + } + return duration, nil +} + +func cloneUploadBaseURL(base *url.URL) *url.URL { + copy := *base + if !strings.HasSuffix(copy.Path, "/") { + copy.Path += "/" + } + return © +} + +func checkedUploadSessionID(location string, createURL *url.URL) (string, error) { + parsed, err := url.Parse(location) + if err != nil || parsed == nil || parsed.RawQuery != "" || parsed.Fragment != "" || parsed.User != nil { + return "", errors.New("invalid upload session location") + } + if parsed.IsAbs() && (parsed.Scheme != createURL.Scheme || !strings.EqualFold(parsed.Host, createURL.Host)) { + return "", errors.New("upload session location changed origin") + } + if parsed.Host != "" && !parsed.IsAbs() { + return "", errors.New("invalid upload session location") + } + match := uploadSessionPath.FindStringSubmatch(parsed.Path) + if len(match) != 2 { + return "", errors.New("invalid upload session path") + } + return match[1], nil +} diff --git a/internal/synapse/pdp_upload_probe_test.go b/internal/synapse/pdp_upload_probe_test.go new file mode 100644 index 0000000..14a650c --- /dev/null +++ b/internal/synapse/pdp_upload_probe_test.go @@ -0,0 +1,131 @@ +package synapse + +import ( + "context" + "errors" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/strahe/synaps3/internal/providerbenchmark" +) + +func TestPDPBatchUploadProbeUploadsSampleWithoutFinalizing(t *testing.T) { + var requests []string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requests = append(requests, r.Method+" "+r.URL.Path) + switch { + case r.Method == http.MethodPost && r.URL.Path == "/proxy/pdp/piece/uploads": + w.Header().Set("Location", "/pdp/piece/uploads/123e4567-e89b-12d3-a456-426614174000") + w.WriteHeader(http.StatusCreated) + case r.Method == http.MethodPut && r.URL.Path == "/proxy/pdp/piece/uploads/123e4567-e89b-12d3-a456-426614174000": + if r.ContentLength != providerbenchmark.SampleBytes { + t.Errorf("content length = %d", r.ContentLength) + } + if r.Header.Get("Content-Type") != "application/octet-stream" { + t.Errorf("content type = %q", r.Header.Get("Content-Type")) + } + count, err := io.Copy(io.Discard, r.Body) + if err != nil { + t.Errorf("read body: %v", err) + } + if int64(count) != providerbenchmark.SampleBytes { + t.Errorf("body bytes = %d", count) + } + w.WriteHeader(http.StatusNoContent) + default: + t.Errorf("unexpected request %s %s", r.Method, r.URL.Path) + w.WriteHeader(http.StatusBadRequest) + } + })) + defer server.Close() + probe := NewPDPBatchUploadProbe(true) + duration, err := probe.Probe(t.Context(), server.URL+"/proxy") + if err != nil || duration <= 0 { + t.Fatalf("Probe = %v, %v", duration, err) + } + if got := strings.Join(requests, ","); got != "POST /proxy/pdp/piece/uploads,PUT /proxy/pdp/piece/uploads/123e4567-e89b-12d3-a456-426614174000" { + t.Fatalf("requests = %s", got) + } +} + +func TestPDPBatchUploadProbeRejectsForeignSessionLocation(t *testing.T) { + var putCalled bool + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodPut { + putCalled = true + } + w.Header().Set("Location", "http://other.example/pdp/piece/uploads/123e4567-e89b-12d3-a456-426614174000") + w.WriteHeader(http.StatusCreated) + })) + defer server.Close() + probe := NewPDPBatchUploadProbe(true) + if _, err := probe.Probe(t.Context(), server.URL); err == nil { + t.Fatal("foreign session location accepted") + } + if putCalled { + t.Fatal("PUT was sent after invalid Location") + } +} + +func TestPDPBatchUploadProbeRejectsFailedPutWithoutFollowingRedirect(t *testing.T) { + var requests []string + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requests = append(requests, r.Method+" "+r.URL.Path) + if r.Method == http.MethodPost { + w.Header().Set("Location", "/pdp/piece/uploads/123e4567-e89b-12d3-a456-426614174000") + w.WriteHeader(http.StatusCreated) + return + } + _, _ = io.Copy(io.Discard, r.Body) + w.Header().Set("Location", "/redirect-target") + w.WriteHeader(http.StatusTemporaryRedirect) + })) + defer server.Close() + probe := NewPDPBatchUploadProbe(true) + if _, err := probe.Probe(t.Context(), server.URL); err == nil || !strings.Contains(err.Error(), "HTTP 307") { + t.Fatalf("redirected PUT = %v", err) + } + if len(requests) != 2 || requests[0] != "POST /pdp/piece/uploads" || requests[1] != "PUT /pdp/piece/uploads/123e4567-e89b-12d3-a456-426614174000" { + t.Fatalf("requests = %v", requests) + } +} + +func TestPDPBatchUploadProbeTimesOutDuringPut(t *testing.T) { + release := make(chan struct{}) + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodPost { + w.Header().Set("Location", "/pdp/piece/uploads/123e4567-e89b-12d3-a456-426614174000") + w.WriteHeader(http.StatusCreated) + return + } + _, _ = io.Copy(io.Discard, r.Body) + select { + case <-r.Context().Done(): + case <-release: + } + })) + defer server.Close() + defer close(release) + probe := NewPDPBatchUploadProbe(true) + probe.timeout = 20 * time.Millisecond + if _, err := probe.Probe(t.Context(), server.URL); !errors.Is(err, context.DeadlineExceeded) { + t.Fatalf("timed out PUT = %v", err) + } +} + +func TestPDPBatchUploadProbeRespectsCallerCancellation(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + cancel() + probe := NewPDPBatchUploadProbe(true) + start := time.Now() + if _, err := probe.Probe(ctx, "http://127.0.0.1:1"); err == nil { + t.Fatal("cancelled context accepted") + } + if time.Since(start) > time.Second { + t.Fatal("cancelled probe took too long") + } +} diff --git a/internal/task/contract.go b/internal/task/contract.go index 2c5b7f7..ef35b7b 100644 --- a/internal/task/contract.go +++ b/internal/task/contract.go @@ -94,6 +94,9 @@ type Definition struct { // CanManualRetry optionally narrows AllowRetry for one failed task based on // its durable failure evidence. Nil preserves the type-wide policy. CanManualRetry func(*model.Task) bool + // OnEngineFailure settles domain state when the engine fails a claim before + // the handler can return its own result. + OnEngineFailure func(*model.Task, string) Settlement } func (d Definition) manualRetryAllowed(task *model.Task) bool { @@ -194,6 +197,7 @@ type Resource string const ( ResourceProviderMutation Resource = "provider_mutation" + ResourceProviderUploadSpeed Resource = "provider_upload_speed" ResourceDestructiveMutation Resource = "destructive_mutation" ResourceWallet Resource = "wallet" ) diff --git a/internal/task/engine.go b/internal/task/engine.go index 34b4032..92752c1 100644 --- a/internal/task/engine.go +++ b/internal/task/engine.go @@ -79,6 +79,7 @@ func NewEngine(config EngineConfig, repos *repository.Repositories, registry *Re config: config, repos: repos, registry: registry, logger: logger, gates: map[Resource]chan struct{}{ ResourceProviderMutation: make(chan struct{}, config.ProviderMutationConcurrency), + ResourceProviderUploadSpeed: make(chan struct{}, 1), ResourceDestructiveMutation: make(chan struct{}, config.DestructiveMutationConcurrency), ResourceWallet: make(chan struct{}, 1), }, @@ -292,7 +293,11 @@ func (e *Engine) executeClaim(parent context.Context, claimed *model.Task) { } func (e *Engine) failClaim(ctx context.Context, claimed *model.Task, reason string, cause error) error { - if err := e.commitResult(ctx, claimed, Fail(cause, reason, nil)); err != nil { + var settlement Settlement + if definition, ok := e.registry.Definition(claimed.Type); ok && definition.OnEngineFailure != nil { + settlement = definition.OnEngineFailure(claimed, reason) + } + if err := e.commitResult(ctx, claimed, Fail(cause, reason, settlement)); err != nil { e.logger.Error("recording task engine failure", "task_id", claimed.ID, "claim_generation", claimed.ClaimGeneration, "error", err) return err } diff --git a/internal/worker/provider_upload_speed_task.go b/internal/worker/provider_upload_speed_task.go new file mode 100644 index 0000000..80b1ec8 --- /dev/null +++ b/internal/worker/provider_upload_speed_task.go @@ -0,0 +1,131 @@ +package worker + +import ( + "context" + "errors" + "time" + + "github.com/strahe/synaps3/internal/db/repository" + "github.com/strahe/synaps3/internal/model" + "github.com/strahe/synaps3/internal/providerbenchmark" + taskengine "github.com/strahe/synaps3/internal/task" + "github.com/strahe/synaps3/internal/types" +) + +func (h *TaskHandlers) providerUploadSpeedHandler() taskengine.Handler { + return taskHandler{ + definition: taskengine.Definition{ + Type: model.TaskTypeProviderUploadSpeedTest, InputVersion: 1, + Codec: taskengine.StrictJSONCodec(func(input *providerbenchmark.Input) error { + if input.ProviderID == "" || len(input.ServiceURLHash) != 64 { + return errors.New("invalid provider upload speed input") + } + _, err := types.ParseOnChainID("provider_id", input.ProviderID) + return err + }), + RetryLimit: new(int), AllowRetry: false, + OnEngineFailure: func(task *model.Task, reason string) taskengine.Settlement { + return func(ctx context.Context, repos *repository.Repositories) error { + return repos.ProviderUploadSpeed.FailActiveTask(ctx, task.ID, reason) + } + }, + }, + execute: h.executeProviderUploadSpeed, + recover: h.recoverProviderUploadSpeed, + } +} + +func (h *TaskHandlers) executeProviderUploadSpeed(ctx context.Context, execution taskengine.Execution) taskengine.Result { + input, err := taskengine.DecodeInput[providerbenchmark.Input](execution) + if err != nil { + return h.failInvalidProviderUploadSpeedInput(execution, err) + } + if h.deps.Observability == nil || h.deps.UploadSpeedProbe == nil { + return h.failProviderUploadSpeed(execution, input, errors.New("upload speed probe unavailable"), "unavailable") + } + id, err := types.ParseOnChainID("provider_id", input.ProviderID) + if err != nil { + return h.failProviderUploadSpeed(execution, input, err, "invalid_input") + } + serviceURL, eligible, err := providerbenchmark.CurrentServiceURL(ctx, h.deps.Observability, id) + if err != nil { + return h.failProviderUploadSpeed(execution, input, err, "unavailable") + } + if !eligible || providerbenchmark.URLHash(serviceURL) != input.ServiceURLHash { + return h.failProviderUploadSpeed(execution, input, errors.New("provider is no longer available at the tested address"), "provider_changed") + } + var duration time.Duration + err = execution.WithResource(ctx, taskengine.ResourceProviderUploadSpeed, func(ctx context.Context) error { + _, effectErr := execution.WithCheckpointedEffect(ctx, taskengine.ResourceProviderMutation, + providerbenchmark.Checkpoint{Attempted: true}, nil, func(ctx context.Context) error { + var probeErr error + duration, probeErr = h.deps.UploadSpeedProbe.Probe(ctx, serviceURL) + return probeErr + }) + return effectErr + }) + if errors.Is(err, taskengine.ErrResourceBusy) { + return taskengine.ResourceWait("Waiting for another speed test or storage operation to finish") + } + if err != nil { + code := "upload_failed" + if errors.Is(err, context.DeadlineExceeded) { + code = "timeout" + } + return h.failProviderUploadSpeed(execution, input, err, code) + } + if duration < time.Millisecond { + duration = time.Millisecond + } + checkpoint := providerbenchmark.Checkpoint{ + Attempted: true, DurationMS: duration.Milliseconds(), + BytesPerSecond: int64(float64(providerbenchmark.SampleBytes) / duration.Seconds()), + } + if err := execution.WriteCheckpoint(ctx, checkpoint); err != nil { + return h.failProviderUploadSpeed(execution, input, err, "record_failed") + } + return h.completeProviderUploadSpeed(execution, input, checkpoint) +} + +func (h *TaskHandlers) recoverProviderUploadSpeed(ctx context.Context, execution taskengine.Execution) taskengine.Result { + input, err := taskengine.DecodeInput[providerbenchmark.Input](execution) + if err != nil { + return h.failInvalidProviderUploadSpeedInput(execution, err) + } + checkpoint, present, err := taskengine.DecodeCheckpoint[providerbenchmark.Checkpoint](execution) + if err != nil { + return h.failProviderUploadSpeed(execution, input, err, "invalid_checkpoint") + } + if present && checkpoint.DurationMS > 0 && checkpoint.BytesPerSecond > 0 { + return h.completeProviderUploadSpeed(execution, input, checkpoint) + } + return h.failProviderUploadSpeed(execution, input, errors.New("upload speed test interrupted"), "interrupted") +} + +func (h *TaskHandlers) completeProviderUploadSpeed(execution taskengine.Execution, input providerbenchmark.Input, checkpoint providerbenchmark.Checkpoint) taskengine.Result { + return taskengine.Complete("Provider upload speed tested", func(ctx context.Context, repos *repository.Repositories) error { + return repos.ProviderUploadSpeed.Finish(ctx, input.ProviderID, execution.ID(), providerbenchmark.StateSucceeded, + checkpoint.DurationMS, checkpoint.BytesPerSecond, "") + }) +} + +func (h *TaskHandlers) failProviderUploadSpeed(execution taskengine.Execution, input providerbenchmark.Input, err error, code string) taskengine.Result { + h.deps.Logger.Warn("provider upload speed test failed", "provider_id", input.ProviderID, "reason", code, "error", err) + message := "Provider upload speed test failed" + if code == "timeout" { + message = "Provider upload speed test timed out" + } + if code == "provider_changed" { + message = "Provider is no longer ready for this test" + } + return taskengine.Fail(errors.New(message), code, func(ctx context.Context, repos *repository.Repositories) error { + return repos.ProviderUploadSpeed.Finish(ctx, input.ProviderID, execution.ID(), providerbenchmark.StateFailed, 0, 0, code) + }) +} + +func (h *TaskHandlers) failInvalidProviderUploadSpeedInput(execution taskengine.Execution, err error) taskengine.Result { + h.deps.Logger.Warn("provider upload speed task input is invalid", "task_id", execution.ID(), "error", err) + return taskengine.Fail(errors.New("upload speed test could not complete"), "invalid_input", func(ctx context.Context, repos *repository.Repositories) error { + return repos.ProviderUploadSpeed.FailActiveTask(ctx, execution.ID(), "invalid_input") + }) +} diff --git a/internal/worker/provider_upload_speed_task_test.go b/internal/worker/provider_upload_speed_task_test.go new file mode 100644 index 0000000..d54885d --- /dev/null +++ b/internal/worker/provider_upload_speed_task_test.go @@ -0,0 +1,167 @@ +package worker_test + +import ( + "context" + "errors" + "sync/atomic" + "testing" + "time" + + "github.com/strahe/synaps3/internal/db/repository" + "github.com/strahe/synaps3/internal/model" + "github.com/strahe/synaps3/internal/observability" + "github.com/strahe/synaps3/internal/providerbenchmark" + taskengine "github.com/strahe/synaps3/internal/task" +) + +type uploadSpeedProbeFunc func(context.Context, string) (time.Duration, error) + +func (f uploadSpeedProbeFunc) Probe(ctx context.Context, url string) (time.Duration, error) { + return f(ctx, url) +} + +func seedSpeedTest(t *testing.T, runtime handlerTestRuntime) *model.Task { + t.Helper() + serviceURL := "https://provider.example" + if err := runtime.repos.Observability.ReplaceProviderStates(t.Context(), time.Now().UTC(), []observability.ProviderState{{ + ProviderID: testOnChainID(t, 101), Status: observability.StatusAvailable, + Active: new(true), HasPDP: new(true), ServiceURL: &serviceURL, LastCheckedAt: time.Now().UTC(), + ReasonCodes: []observability.ReasonCode{}, Evidence: map[string]any{}, + }}); err != nil { + t.Fatal(err) + } + input := providerbenchmark.Input{ProviderID: "101", ServiceURLHash: providerbenchmark.URLHash(serviceURL)} + taskRow, _, err := runtime.service.EnqueueTx(t.Context(), taskengine.EnqueueRequest{ + Type: model.TaskTypeProviderUploadSpeedTest, IdempotencyKey: "speed-101", Input: input, + SubjectType: "provider", SubjectKey: "101", + }, func(ctx context.Context, repos *repository.Repositories, taskRow *model.Task, _ bool) error { + return repos.ProviderUploadSpeed.Begin(ctx, "101", input.ServiceURLHash, taskRow.ID) + }) + if err != nil { + t.Fatal(err) + } + return taskRow +} + +func TestProviderUploadSpeedTaskRecordsSuccessfulMeasurement(t *testing.T) { + var calls atomic.Int32 + runtime := newHandlerTestRuntime(t, handlerRuntimeOptions{uploadSpeedProbe: uploadSpeedProbeFunc(func(_ context.Context, url string) (time.Duration, error) { + calls.Add(1) + if url != "https://provider.example" { + t.Errorf("probe URL = %q", url) + } + return 2 * time.Second, nil + })}) + taskRow := seedSpeedTest(t, runtime) + cancel, done := runHandlerEngine(t, runtime) + defer stopHandlerEngine(t, cancel, done) + waitForTask(t, runtime.repos, taskRow.ID, func(task *model.Task) bool { return task.Status == model.TaskStatusCompleted }) + row, err := runtime.repos.ProviderUploadSpeed.Get(t.Context(), "101") + if err != nil || row == nil || row.State != providerbenchmark.StateSucceeded || row.DurationMS == nil || *row.DurationMS != 2000 || row.ActiveTaskID != nil { + t.Fatalf("speed result = %#v, %v", row, err) + } + if calls.Load() != 1 { + t.Fatalf("probe calls = %d", calls.Load()) + } +} + +func TestProviderUploadSpeedEngineFailureReleasesActiveTest(t *testing.T) { + for _, tt := range []struct { + name string + reason string + breakTask func(*testing.T, handlerTestRuntime, *model.Task) + }{ + { + name: "handler panic", reason: "handler_panic", + breakTask: func(*testing.T, handlerTestRuntime, *model.Task) {}, + }, + { + name: "invalid input hash", reason: "invalid_input_hash", + breakTask: func(t *testing.T, runtime handlerTestRuntime, taskRow *model.Task) { + if _, err := runtime.db.NewUpdate().Model((*model.Task)(nil)).Set("input_hash = ?", providerbenchmark.URLHash("wrong input")). + Where("id = ?", taskRow.ID).Exec(t.Context()); err != nil { + t.Fatal(err) + } + }, + }, + } { + t.Run(tt.name, func(t *testing.T) { + var calls atomic.Int32 + runtime := newHandlerTestRuntime(t, handlerRuntimeOptions{uploadSpeedProbe: uploadSpeedProbeFunc(func(context.Context, string) (time.Duration, error) { + calls.Add(1) + panic("probe failed unexpectedly") + })}) + taskRow := seedSpeedTest(t, runtime) + tt.breakTask(t, runtime, taskRow) + cancel, done := runHandlerEngine(t, runtime) + defer stopHandlerEngine(t, cancel, done) + failed := waitForTask(t, runtime.repos, taskRow.ID, func(task *model.Task) bool { return task.Status == model.TaskStatusFailed }) + if failed.FailureReason == nil || *failed.FailureReason != tt.reason { + t.Fatalf("task failure reason = %v, want %s", failed.FailureReason, tt.reason) + } + row, err := runtime.repos.ProviderUploadSpeed.Get(t.Context(), "101") + if err != nil || row == nil || row.State != providerbenchmark.StateFailed || row.ActiveTaskID != nil || + row.FailureCode == nil || *row.FailureCode != tt.reason { + t.Fatalf("speed result after engine failure = %#v, err=%v", row, err) + } + if tt.reason == "handler_panic" && calls.Load() != 1 || tt.reason == "invalid_input_hash" && calls.Load() != 0 { + t.Fatalf("probe calls = %d, reason=%s", calls.Load(), tt.reason) + } + }) + } +} + +func TestProviderUploadSpeedRecoveryNeverReuploads(t *testing.T) { + var calls atomic.Int32 + runtime := newHandlerTestRuntime(t, handlerRuntimeOptions{uploadSpeedProbe: uploadSpeedProbeFunc(func(context.Context, string) (time.Duration, error) { + calls.Add(1) + return 0, errors.New("should not upload") + })}) + taskRow := seedSpeedTest(t, runtime) + if _, err := runtime.db.NewUpdate().Model((*model.Task)(nil)).Set("resume_mode = ?", model.TaskResumeModeRecover). + Where("id = ?", taskRow.ID).Exec(t.Context()); err != nil { + t.Fatal(err) + } + if _, err := runtime.db.NewUpdate().TableExpr("task_payloads").Set("checkpoint_json = ?", `{"attempted":true}`). + Where("task_id = ?", taskRow.ID).Exec(t.Context()); err != nil { + t.Fatal(err) + } + cancel, done := runHandlerEngine(t, runtime) + defer stopHandlerEngine(t, cancel, done) + waitForTask(t, runtime.repos, taskRow.ID, func(task *model.Task) bool { return task.Status == model.TaskStatusFailed }) + row, err := runtime.repos.ProviderUploadSpeed.Get(t.Context(), "101") + if err != nil || row == nil || row.State != providerbenchmark.StateFailed || row.ActiveTaskID != nil { + t.Fatalf("recovered result = %#v, %v", row, err) + } + if calls.Load() != 0 { + t.Fatalf("probe calls during recovery = %d", calls.Load()) + } +} + +func TestProviderUploadSpeedRecoverySettlesCompletedPutWithoutReupload(t *testing.T) { + var calls atomic.Int32 + runtime := newHandlerTestRuntime(t, handlerRuntimeOptions{uploadSpeedProbe: uploadSpeedProbeFunc(func(context.Context, string) (time.Duration, error) { + calls.Add(1) + return 0, errors.New("should not upload") + })}) + taskRow := seedSpeedTest(t, runtime) + if _, err := runtime.db.NewUpdate().Model((*model.Task)(nil)).Set("resume_mode = ?", model.TaskResumeModeRecover). + Where("id = ?", taskRow.ID).Exec(t.Context()); err != nil { + t.Fatal(err) + } + if _, err := runtime.db.NewUpdate().TableExpr("task_payloads"). + Set("checkpoint_json = ?", `{"attempted":true,"duration_ms":2000,"bytes_per_second":16777216}`). + Where("task_id = ?", taskRow.ID).Exec(t.Context()); err != nil { + t.Fatal(err) + } + cancel, done := runHandlerEngine(t, runtime) + defer stopHandlerEngine(t, cancel, done) + waitForTask(t, runtime.repos, taskRow.ID, func(task *model.Task) bool { return task.Status == model.TaskStatusCompleted }) + row, err := runtime.repos.ProviderUploadSpeed.Get(t.Context(), "101") + if err != nil || row == nil || row.State != providerbenchmark.StateSucceeded || row.BytesPerSecond == nil || *row.BytesPerSecond != 16777216 || row.ActiveTaskID != nil { + t.Fatalf("recovered result = %#v, %v", row, err) + } + if calls.Load() != 0 { + t.Fatalf("probe calls during recovery = %d", calls.Load()) + } +} diff --git a/internal/worker/task_handlers.go b/internal/worker/task_handlers.go index 99ec45d..55641a0 100644 --- a/internal/worker/task_handlers.go +++ b/internal/worker/task_handlers.go @@ -32,14 +32,17 @@ type TaskHandlerDependencies struct { Terminator synapse.ServiceTerminator Epochs synapse.ChainEpochReader Observability *observability.Service - ParkedPieces synapse.ParkedPieceChecker - EvictionPolicy cache.EvictionPolicy - MaxCacheBytes int64 - LRUHighPercent int - LRULowPercent int - DefaultCopies int - MaxRetries int - Logger *slog.Logger + UploadSpeedProbe interface { + Probe(context.Context, string) (time.Duration, error) + } + ParkedPieces synapse.ParkedPieceChecker + EvictionPolicy cache.EvictionPolicy + MaxCacheBytes int64 + LRUHighPercent int + LRULowPercent int + DefaultCopies int + MaxRetries int + Logger *slog.Logger } type EventPublisher interface { @@ -97,6 +100,7 @@ func (h *TaskHandlers) RegisterCore(registry *taskengine.Registry) error { h.storageCleanupHandler(), h.walletHandler(), h.observabilityHandler(), + h.providerUploadSpeedHandler(), h.gcHandler(), } { if err := registry.Register(handler); err != nil { diff --git a/internal/worker/task_handlers_test.go b/internal/worker/task_handlers_test.go index c727550..4e425c7 100644 --- a/internal/worker/task_handlers_test.go +++ b/internal/worker/task_handlers_test.go @@ -27,6 +27,7 @@ import ( "github.com/strahe/synaps3/internal/cacheeviction" "github.com/strahe/synaps3/internal/db/repository" "github.com/strahe/synaps3/internal/model" + "github.com/strahe/synaps3/internal/observability" "github.com/strahe/synaps3/internal/storagecleanup" "github.com/strahe/synaps3/internal/storagecommit" "github.com/strahe/synaps3/internal/storagepipeline" @@ -70,14 +71,17 @@ type handlerRuntimeOptions struct { terminator synapse.ServiceTerminator epochs synapse.ChainEpochReader parkedPieces synapse.ParkedPieceChecker - policy cache.EvictionPolicy - maxBytes int64 - highPercent int - lowPercent int - concurrency int - maxRetries *int - leaseDuration time.Duration - register func(*worker.TaskHandlers, *taskengine.Registry) error + uploadSpeedProbe interface { + Probe(context.Context, string) (time.Duration, error) + } + policy cache.EvictionPolicy + maxBytes int64 + highPercent int + lowPercent int + concurrency int + maxRetries *int + leaseDuration time.Duration + register func(*worker.TaskHandlers, *taskengine.Registry) error } func newHandlerTestRuntime(t *testing.T, options handlerRuntimeOptions) handlerTestRuntime { @@ -98,13 +102,18 @@ func newHandlerTestRuntime(t *testing.T, options handlerRuntimeOptions) handlerT if options.maxRetries != nil { maxRetries = *options.maxRetries } + var observabilityService *observability.Service + if options.uploadSpeedProbe != nil { + observabilityService = observability.NewService(observability.ServiceOptions{Store: repos.Observability}) + } handlers, err := worker.NewTaskHandlers(worker.TaskHandlerDependencies{ Repositories: repos, Events: options.events, Cache: cacheStore, CacheGate: gate, CacheTracker: tracker, Storage: storageClient, Wallet: options.wallet, Receipts: options.receipts, WalletBroadcastTimeout: options.walletBroadcastTimeout, WalletReceiptTimeout: options.walletReceiptTimeout, Terminator: options.terminator, Epochs: options.epochs, - ParkedPieces: options.parkedPieces, + ParkedPieces: options.parkedPieces, + Observability: observabilityService, UploadSpeedProbe: options.uploadSpeedProbe, EvictionPolicy: options.policy, MaxCacheBytes: options.maxBytes, LRUHighPercent: options.highPercent, LRULowPercent: options.lowPercent, DefaultCopies: 2, MaxRetries: maxRetries, Logger: slog.Default(), diff --git a/ui/src/api/client.ts b/ui/src/api/client.ts index 5317309..ee6db0e 100644 --- a/ui/src/api/client.ts +++ b/ui/src/api/client.ts @@ -157,6 +157,16 @@ export interface ObservabilityProviderFacts { export interface ObservabilityProviderObservation { facts: ObservabilityProviderFacts signal: ObservabilitySignal + upload_speed_test?: ProviderUploadSpeedTest +} + +export interface ProviderUploadSpeedTest { + state: 'testing' | 'succeeded' | 'failed' | 'stale' + sample_bytes: number + duration_ms?: number + bytes_per_second?: number + tested_at?: string + failure_code?: string } export interface ObservabilityDataSetFacts { @@ -1164,6 +1174,11 @@ export const api = { fetchJSON>( `/observability/providers${observabilityListQuery(params)}` ), + testProviderUploadSpeed: (providerID: string) => + fetchJSON<{ task_id: number; state: 'testing' }>( + `/observability/providers/${encodeURIComponent(providerID)}/upload-speed-test`, + { method: 'POST' } + ), getObservabilityDataSets: (params: ObservabilityListParams = {}) => fetchJSON>( `/observability/data-sets${observabilityListQuery(params)}` diff --git a/ui/src/components/storage-topology/StorageTopologyDetailSheet.tsx b/ui/src/components/storage-topology/StorageTopologyDetailSheet.tsx index 8de6154..83a6840 100644 --- a/ui/src/components/storage-topology/StorageTopologyDetailSheet.tsx +++ b/ui/src/components/storage-topology/StorageTopologyDetailSheet.tsx @@ -12,6 +12,7 @@ import { Button } from '@/components/ui/button' import { ScrollArea } from '@/components/ui/scroll-area' import { Sheet, SheetContent, SheetDescription, SheetHeader, SheetTitle } from '@/components/ui/sheet' import { activePiecesValue } from '@/lib/data-set-storage-health' +import { providerUploadSampleSize, providerUploadSpeedLabel } from '@/lib/provider-upload-speed' import { replicaLabel } from '@/lib/storage-status-labels' import { dataSetDisplayLabel, @@ -25,7 +26,7 @@ import { relatedDataSetsForProviderNode, type StorageTopologyGraph, } from '@/lib/storage-topology' -import { cn, formatNumber } from '@/lib/utils' +import { cn, formatNumber, timeAgo } from '@/lib/utils' export function TopologyDetailSheet({ selection, @@ -275,6 +276,14 @@ function ProviderDetailContent({ provider }: { provider: ObservabilityProviderOb /> + + + + + ) } diff --git a/ui/src/components/storage-topology/StorageTopologyTables.tsx b/ui/src/components/storage-topology/StorageTopologyTables.tsx index f56971d..eebdc22 100644 --- a/ui/src/components/storage-topology/StorageTopologyTables.tsx +++ b/ui/src/components/storage-topology/StorageTopologyTables.tsx @@ -1,5 +1,5 @@ import { Link } from '@tanstack/react-router' -import { Database, Info } from 'lucide-react' +import { Database, Gauge, Info, Loader2 } from 'lucide-react' import type { ReactNode } from 'react' import type { ObservabilityDataSetObservation, ObservabilityProviderObservation } from '@/api/client' import { CopyableValue } from '@/components/app/CopyableValue' @@ -9,7 +9,9 @@ import { Empty, EmptyDescription, EmptyHeader, EmptyMedia, EmptyTitle } from '@/ import { Pagination, PaginationContent, PaginationItem } from '@/components/ui/pagination' import { Skeleton } from '@/components/ui/skeleton' import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from '@/components/ui/table' +import { Tooltip, TooltipContent, TooltipTrigger } from '@/components/ui/tooltip' import { activePiecesValue } from '@/lib/data-set-storage-health' +import { canTestProviderUploadSpeed, providerUploadSpeedLabel } from '@/lib/provider-upload-speed' import { replicaLabel } from '@/lib/storage-status-labels' import { dataSetChainIDValue, @@ -32,6 +34,8 @@ export function ProvidersTableCard({ contextNote, onPageChange, onSelect, + onTestUploadSpeed, + testingProviderID, }: { rows: StorageTopologyProviderRow[] total: number @@ -42,6 +46,8 @@ export function ProvidersTableCard({ contextNote?: string onPageChange: (page: number) => void onSelect: (item: StorageTopologyProviderRow) => void + onTestUploadSpeed: (providerID: string) => void + testingProviderID?: string }) { return (
@@ -69,7 +75,30 @@ export function ProvidersTableCard({ - {row.status} +
+ {row.status} + {canTestProviderUploadSpeed(row.provider) && ( + + + + + Test upload speed with 32 MiB + + )} + {row.provider?.upload_speed_test && ( + + {providerUploadSpeedLabel(row.provider.upload_speed_test)} + + )} +
{row.provider ? providerFactsSummary(row.provider) : '—'} diff --git a/ui/src/hooks/queries.ts b/ui/src/hooks/queries.ts index d3fd500..798e642 100644 --- a/ui/src/hooks/queries.ts +++ b/ui/src/hooks/queries.ts @@ -370,6 +370,17 @@ export function useObservabilityProviders( }) } +export function useTestProviderUploadSpeed() { + const qc = useQueryClient() + return useMutation({ + mutationFn: api.testProviderUploadSpeed, + onSuccess: () => { + qc.invalidateQueries({ queryKey: ['observabilityProviders'] }) + qc.invalidateQueries({ queryKey: ['tasks'] }) + }, + }) +} + export function useObservabilityDataSets(params: ObservabilityListParams, enabled = true) { return useQuery({ queryKey: ['observabilityDataSets', params], diff --git a/ui/src/lib/overview.ts b/ui/src/lib/overview.ts index 7f3d3f8..fe5bd2d 100644 --- a/ui/src/lib/overview.ts +++ b/ui/src/lib/overview.ts @@ -91,6 +91,7 @@ const taskOperationLabels: Record = { storage_dataset_retire: 'Retire service', wallet_operation: 'Wallet request', observability_refresh: 'Refresh health', + provider_upload_speed_test: 'Test provider upload speed', task_gc: 'Remove expired task records', } diff --git a/ui/src/lib/provider-upload-speed.ts b/ui/src/lib/provider-upload-speed.ts new file mode 100644 index 0000000..630434b --- /dev/null +++ b/ui/src/lib/provider-upload-speed.ts @@ -0,0 +1,45 @@ +import type { ObservabilityProviderObservation, ProviderUploadSpeedTest } from '../api/client.ts' + +export function canTestProviderUploadSpeed(provider?: ObservabilityProviderObservation) { + return Boolean( + provider?.signal.status === 'available' && + !provider.signal.freshness.stale && + provider.facts.active && + provider.facts.has_pdp && + provider.facts.service_url && + provider.upload_speed_test?.state !== 'testing' + ) +} + +export function providerUploadSpeedLabel(test?: ProviderUploadSpeedTest): string { + if (!test) return 'Not tested' + switch (test.state) { + case 'testing': + return 'Testing upload speed…' + case 'stale': + return 'Outdated — Service URL no longer matches' + case 'failed': + switch (test.failure_code) { + case 'timeout': + return 'Upload test timed out' + case 'interrupted': + return 'Upload test interrupted — test again' + case 'unavailable': + return 'Upload test could not run' + case 'provider_changed': + return 'Provider no longer ready for this test' + case 'record_failed': + return 'Upload speed not recorded' + default: + return 'Upload test failed' + } + case 'succeeded': + if (!test.bytes_per_second) return 'Upload speed not recorded' + return `${(test.bytes_per_second / (1024 * 1024)).toFixed(1)} MiB/s` + } +} + +export function providerUploadSampleSize(test?: ProviderUploadSpeedTest): string { + if (!test) return '—' + return `${Number((test.sample_bytes / (1024 * 1024)).toFixed(1))} MiB` +} diff --git a/ui/src/routes/storage-topology.tsx b/ui/src/routes/storage-topology.tsx index 37c2148..b1e7fdb 100644 --- a/ui/src/routes/storage-topology.tsx +++ b/ui/src/routes/storage-topology.tsx @@ -2,7 +2,7 @@ import { useQueryClient } from '@tanstack/react-query' import { createFileRoute, useNavigate } from '@tanstack/react-router' import { Database, RefreshCw, TriangleAlert } from 'lucide-react' import { lazy, Suspense, useEffect, useMemo, useState } from 'react' -import type { ObservabilityDataSetObservation, ObservabilityProviderObservation } from '@/api/client' +import { APIError, type ObservabilityDataSetObservation, type ObservabilityProviderObservation } from '@/api/client' import { PageHeader } from '@/components/app/PageHeader' import { TopologyDetailSheet } from '@/components/storage-topology/StorageTopologyDetailSheet' import { DataSetsTableCard, ProvidersTableCard } from '@/components/storage-topology/StorageTopologyTables' @@ -11,7 +11,7 @@ import { Alert, AlertDescription, AlertTitle } from '@/components/ui/alert' import { Button } from '@/components/ui/button' import { Empty, EmptyDescription, EmptyHeader, EmptyMedia, EmptyTitle } from '@/components/ui/empty' import { Skeleton } from '@/components/ui/skeleton' -import { useObservabilityDataSets, useObservabilityProviders } from '@/hooks/queries' +import { useObservabilityDataSets, useObservabilityProviders, useTestProviderUploadSpeed } from '@/hooks/queries' import { buildStorageTopologyGraph, buildTopologyProviderOptions, @@ -117,6 +117,8 @@ function StorageTopologyPage() { const [dataSetPage, setDataSetPage] = useState(1) const [selection, setSelection] = useState(null) const [pinnedContext, setPinnedContext] = useState(null) + const [uploadSpeedTestError, setUploadSpeedTestError] = useState(null) + const testProviderUploadSpeed = useTestProviderUploadSpeed() const selectionProviderID = search.selection_provider ?? observabilityProviderParam(filters.provider) const selectionBucketName = search.selection_bucket ?? observabilityBucketParam(filters.bucket) @@ -430,6 +432,18 @@ function StorageTopologyPage() { qc.invalidateQueries({ queryKey: ['observabilityDataSets'] }) } + function startProviderUploadSpeedTest(providerID: string) { + setUploadSpeedTestError(null) + testProviderUploadSpeed.mutate(providerID, { + onError: (error) => + setUploadSpeedTestError( + error instanceof APIError && error.status === 409 + ? `A test is already running, or provider #${providerID} is not available for testing. Check its details before trying again.` + : `Could not start the upload speed test for provider #${providerID}. Try again.` + ), + }) + } + const pageClassName = tab === 'topology' ? 'flex h-[calc(100svh-3.5rem)] min-h-0 min-w-0 flex-col gap-4 overflow-hidden px-6 pt-6 pb-0 md:h-svh' @@ -468,6 +482,13 @@ function StorageTopologyPage() {
)} + {tab === 'providers' && uploadSpeedTestError && ( + + Upload speed test not started + {uploadSpeedTestError} + + )} + {snapshotLoading ? ( ) : snapshotError ? ( @@ -495,6 +516,8 @@ function StorageTopologyPage() { } onPageChange={setProviderPage} onSelect={selectProvider} + onTestUploadSpeed={startProviderUploadSpeedTest} + testingProviderID={testProviderUploadSpeed.isPending ? testProviderUploadSpeed.variables : undefined} /> ) : ( { + const originalFetch = globalThis.fetch + let url = '' + let method = '' + let csrf = '' + globalThis.fetch = (async (input, init) => { + url = input.toString() + method = init?.method ?? '' + csrf = new Headers(init?.headers).get('X-SynapS3-CSRF') ?? '' + return new Response(JSON.stringify({ task_id: 42, state: 'testing' }), { + status: 202, + headers: { 'Content-Type': 'application/json' }, + }) + }) as typeof fetch + try { + setAdminCSRFToken('speed-csrf') + const result = await api.testProviderUploadSpeed('101') + assert.deepEqual(result, { task_id: 42, state: 'testing' }) + } finally { + setAdminCSRFToken('') + globalThis.fetch = originalFetch + } + assert.equal(url, '/api/v1/observability/providers/101/upload-speed-test') + assert.equal(method, 'POST') + assert.equal(csrf, 'speed-csrf') +}) + test('object download URL encodes bucket name and object key', () => { assert.equal( api.getObjectDownloadUrl('bucket-a', 'reports/April summary.txt'), diff --git a/ui/test/provider-upload-speed.test.ts b/ui/test/provider-upload-speed.test.ts new file mode 100644 index 0000000..62f8c09 --- /dev/null +++ b/ui/test/provider-upload-speed.test.ts @@ -0,0 +1,65 @@ +import assert from 'node:assert/strict' +import test from 'node:test' + +import type { ObservabilityProviderObservation } from '../src/api/client.ts' +import { + canTestProviderUploadSpeed, + providerUploadSampleSize, + providerUploadSpeedLabel, +} from '../src/lib/provider-upload-speed.ts' + +const available: ObservabilityProviderObservation = { + facts: { provider_id: '101', active: true, has_pdp: true, service_url: 'https://provider.example' }, + signal: { status: 'available', level: 'ok', reason_codes: [], freshness: { stale: false, warnings: [] } }, +} + +test('upload speed action requires a fresh available provider and no active test', () => { + assert.equal(canTestProviderUploadSpeed(available), true) + assert.equal( + canTestProviderUploadSpeed({ ...available, upload_speed_test: { state: 'testing', sample_bytes: 32 << 20 } }), + false + ) + assert.equal( + canTestProviderUploadSpeed({ + ...available, + signal: { ...available.signal, freshness: { stale: true, warnings: [] } }, + }), + false + ) + assert.equal(canTestProviderUploadSpeed({ ...available, facts: { ...available.facts, has_pdp: false } }), false) + assert.equal( + canTestProviderUploadSpeed({ ...available, upload_speed_test: { state: 'stale', sample_bytes: 32 << 20 } }), + true + ) +}) + +test('upload speed labels distinguish results from failed or outdated tests', () => { + assert.equal( + providerUploadSpeedLabel({ state: 'succeeded', sample_bytes: 32 << 20, bytes_per_second: 10 << 20 }), + '10.0 MiB/s' + ) + assert.equal( + providerUploadSpeedLabel({ state: 'stale', sample_bytes: 32 << 20 }), + 'Outdated — Service URL no longer matches' + ) + for (const [failureCode, label] of [ + ['timeout', 'Upload test timed out'], + ['interrupted', 'Upload test interrupted — test again'], + ['unavailable', 'Upload test could not run'], + ['provider_changed', 'Provider no longer ready for this test'], + ['record_failed', 'Upload speed not recorded'], + ['upload_failed', 'Upload test failed'], + ]) { + assert.equal( + providerUploadSpeedLabel({ state: 'failed', sample_bytes: 32 << 20, failure_code: failureCode }), + label + ) + } + assert.equal(providerUploadSpeedLabel({ state: 'succeeded', sample_bytes: 32 << 20 }), 'Upload speed not recorded') +}) + +test('sample size describes the saved test rather than an unstarted test', () => { + assert.equal(providerUploadSampleSize(), '—') + assert.equal(providerUploadSampleSize({ state: 'succeeded', sample_bytes: 64 << 20 }), '64 MiB') + assert.equal(providerUploadSampleSize({ state: 'testing', sample_bytes: 3 << 19 }), '1.5 MiB') +})