package socket import ( "context" "errors" "net" "testing" "time" altv1 "git.toki-labs.com/toki/alt/packages/contracts/gen/go/alt/v1" "git.toki-labs.com/toki/alt/packages/domain/backtest" "git.toki-labs.com/toki/alt/services/worker/internal/config" workerContracts "git.toki-labs.com/toki/alt/services/worker/internal/contracts" "git.toki-labs.com/toki/alt/services/worker/internal/storage" protoSocket "git.toki-labs.com/toki/proto-socket/go" ) type fakeStarter struct { gotSpec backtest.RunSpec run backtest.Run err error } func (f *fakeStarter) StartBacktest(ctx context.Context, spec backtest.RunSpec) (backtest.Run, error) { f.gotSpec = spec if f.err != nil { return backtest.Run{}, f.err } return f.run, nil } type fakeAnalysisStore struct { listStatus backtest.RunStatus runs []backtest.Run detail storage.RunDetail compareIDs []backtest.RunID results []backtest.Result err error } func (f *fakeAnalysisStore) ListRuns(ctx context.Context, status backtest.RunStatus) ([]backtest.Run, error) { f.listStatus = status return f.runs, f.err } func (f *fakeAnalysisStore) GetRunDetail(ctx context.Context, id backtest.RunID) (storage.RunDetail, error) { if f.err != nil { return storage.RunDetail{}, f.err } return f.detail, nil } func (f *fakeAnalysisStore) CompareResults(ctx context.Context, ids []backtest.RunID) ([]backtest.Result, error) { f.compareIDs = ids return f.results, f.err } type fakeResultStore struct { result backtest.Result err error } func (f *fakeResultStore) UpsertResult(ctx context.Context, result backtest.Result) error { return nil } func (f *fakeResultStore) GetResult(ctx context.Context, id backtest.RunID) (backtest.Result, error) { if f.err != nil { return backtest.Result{}, f.err } return f.result, nil } func validStartRequest() *altv1.StartBacktestRequest { return &altv1.StartBacktestRequest{ Spec: &altv1.BacktestRunSpec{ StrategyId: "strat-abc", Market: altv1.Market_MARKET_KR, Timeframe: altv1.Timeframe_TIMEFRAME_DAILY, FromUnixMs: time.Date(2026, 5, 1, 0, 0, 0, 0, time.UTC).UnixMilli(), ToUnixMs: time.Date(2026, 5, 15, 0, 0, 0, 0, time.UTC).UnixMilli(), }, } } func TestHandleStartBacktestSuccess(t *testing.T) { starter := &fakeStarter{run: backtest.Run{ID: "run-1", Status: backtest.RunStatusPending}} deps := BacktestDeps{Starter: starter} resp, err := handleStartBacktest(deps, validStartRequest()) if err != nil { t.Fatalf("unexpected error: %v", err) } if resp.GetRun().GetId() != "run-1" { t.Errorf("expected run id run-1, got %q", resp.GetRun().GetId()) } if resp.GetRun().GetStatus() != altv1.BacktestRunStatus_BACKTEST_RUN_STATUS_PENDING { t.Errorf("expected pending status, got %v", resp.GetRun().GetStatus()) } if starter.gotSpec.StrategyID != "strat-abc" { t.Errorf("starter received wrong spec: %+v", starter.gotSpec) } } func TestHandleStartBacktestInvalidSpec(t *testing.T) { deps := BacktestDeps{Starter: &fakeStarter{}} resp, err := handleStartBacktest(deps, &altv1.StartBacktestRequest{}) if err != nil { t.Fatalf("expected typed validation response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInvalidRequest) } func TestHandleStartBacktestUnavailable(t *testing.T) { resp, err := handleStartBacktest(BacktestDeps{}, validStartRequest()) if err != nil { t.Fatalf("expected typed unavailable response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorUnavailable) } func TestHandleStartBacktestStarterError(t *testing.T) { deps := BacktestDeps{Starter: &fakeStarter{err: errors.New("boom")}} resp, err := handleStartBacktest(deps, validStartRequest()) if err != nil { t.Fatalf("expected typed internal response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInternal) } func TestHandleListBacktestRunsForwardsStatusFilter(t *testing.T) { store := &fakeAnalysisStore{runs: []backtest.Run{{ID: "run-1"}, {ID: "run-2"}}} deps := BacktestDeps{Analysis: store} resp, err := handleListBacktestRuns(deps, &altv1.ListBacktestRunsRequest{ Status: altv1.BacktestRunStatus_BACKTEST_RUN_STATUS_SUCCEEDED, }) if err != nil { t.Fatalf("unexpected error: %v", err) } if store.listStatus != backtest.RunStatusSucceeded { t.Errorf("expected status filter succeeded, got %q", store.listStatus) } if len(resp.GetRuns()) != 2 { t.Errorf("expected 2 runs, got %d", len(resp.GetRuns())) } } func TestHandleListBacktestRunsEmpty(t *testing.T) { deps := BacktestDeps{Analysis: &fakeAnalysisStore{}} resp, err := handleListBacktestRuns(deps, &altv1.ListBacktestRunsRequest{}) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(resp.GetRuns()) != 0 { t.Errorf("expected empty runs, got %d", len(resp.GetRuns())) } } func TestHandleGetBacktestRunSuccess(t *testing.T) { store := &fakeAnalysisStore{detail: storage.RunDetail{ Run: backtest.Run{ID: "run-1", Status: backtest.RunStatusSucceeded}, }} deps := BacktestDeps{Analysis: store} resp, err := handleGetBacktestRun(deps, &altv1.GetBacktestRunRequest{RunId: "run-1"}) if err != nil { t.Fatalf("unexpected error: %v", err) } if resp.GetRun().GetId() != "run-1" { t.Errorf("run id mismatch: %q", resp.GetRun().GetId()) } if resp.GetRun().GetStatus() != altv1.BacktestRunStatus_BACKTEST_RUN_STATUS_SUCCEEDED { t.Errorf("expected succeeded status, got %v", resp.GetRun().GetStatus()) } } func TestHandleGetBacktestRunNotFound(t *testing.T) { store := &fakeAnalysisStore{err: storage.ErrRunNotFound} deps := BacktestDeps{Analysis: store} resp, err := handleGetBacktestRun(deps, &altv1.GetBacktestRunRequest{RunId: "missing"}) if err != nil { t.Fatalf("expected typed not_found response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorNotFound) } func TestHandleGetBacktestRunRequiresRunID(t *testing.T) { deps := BacktestDeps{Analysis: &fakeAnalysisStore{}} resp, err := handleGetBacktestRun(deps, &altv1.GetBacktestRunRequest{}) if err != nil { t.Fatalf("expected typed validation response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInvalidRequest) } func TestHandleGetBacktestRunUnavailable(t *testing.T) { resp, err := handleGetBacktestRun(BacktestDeps{}, &altv1.GetBacktestRunRequest{RunId: "run-1"}) if err != nil { t.Fatalf("expected typed unavailable response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorUnavailable) } func TestHandleGetBacktestRunDetailWithResult(t *testing.T) { store := &fakeAnalysisStore{detail: storage.RunDetail{ Run: backtest.Run{ID: "run-1", Status: backtest.RunStatusSucceeded}, Result: backtest.Result{RunID: "run-1"}, HasResult: true, }} deps := BacktestDeps{Analysis: store} resp, err := handleGetBacktestRunDetail(deps, &altv1.GetBacktestRunDetailRequest{RunId: "run-1"}) if err != nil { t.Fatalf("unexpected error: %v", err) } if resp.GetRun().GetId() != "run-1" { t.Errorf("run id mismatch: %q", resp.GetRun().GetId()) } if resp.GetResult() == nil { t.Error("expected result to be present when HasResult is true") } } func TestHandleGetBacktestRunDetailWithoutResult(t *testing.T) { store := &fakeAnalysisStore{detail: storage.RunDetail{ Run: backtest.Run{ID: "run-1", Status: backtest.RunStatusRunning}, HasResult: false, }} deps := BacktestDeps{Analysis: store} resp, err := handleGetBacktestRunDetail(deps, &altv1.GetBacktestRunDetailRequest{RunId: "run-1"}) if err != nil { t.Fatalf("unexpected error: %v", err) } if resp.GetResult() != nil { t.Error("expected nil result when run has no result yet") } } func TestHandleGetBacktestRunDetailNotFound(t *testing.T) { store := &fakeAnalysisStore{err: storage.ErrRunNotFound} deps := BacktestDeps{Analysis: store} resp, err := handleGetBacktestRunDetail(deps, &altv1.GetBacktestRunDetailRequest{RunId: "missing"}) if err != nil { t.Fatalf("expected typed not_found response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorNotFound) } func TestHandleGetBacktestRunDetailRequiresRunID(t *testing.T) { deps := BacktestDeps{Analysis: &fakeAnalysisStore{}} resp, err := handleGetBacktestRunDetail(deps, &altv1.GetBacktestRunDetailRequest{}) if err != nil { t.Fatalf("expected typed validation response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInvalidRequest) } func TestHandleGetBacktestResultSuccess(t *testing.T) { store := &fakeResultStore{result: backtest.Result{RunID: "run-1"}} deps := BacktestDeps{Results: store} resp, err := handleGetBacktestResult(deps, &altv1.GetBacktestResultRequest{RunId: "run-1"}) if err != nil { t.Fatalf("unexpected error: %v", err) } if resp.GetResult().GetRunId() != "run-1" { t.Errorf("result run id mismatch: %q", resp.GetResult().GetRunId()) } } func TestHandleGetBacktestResultNotFound(t *testing.T) { store := &fakeResultStore{err: storage.ErrResultNotFound} deps := BacktestDeps{Results: store} resp, err := handleGetBacktestResult(deps, &altv1.GetBacktestResultRequest{RunId: "run-1"}) if err != nil { t.Fatalf("expected typed not_found response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorNotFound) } func TestHandleCompareBacktestRunsSuccess(t *testing.T) { store := &fakeAnalysisStore{results: []backtest.Result{{RunID: "run-1"}, {RunID: "run-2"}}} deps := BacktestDeps{Analysis: store} resp, err := handleCompareBacktestRuns(deps, &altv1.CompareBacktestRunsRequest{RunIds: []string{"run-1", "run-2"}}) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(store.compareIDs) != 2 { t.Errorf("expected 2 ids forwarded, got %d", len(store.compareIDs)) } if len(resp.GetResults()) != 2 { t.Errorf("expected 2 results, got %d", len(resp.GetResults())) } } func TestHandleCompareBacktestRunsEmptyIsNoOp(t *testing.T) { deps := BacktestDeps{} resp, err := handleCompareBacktestRuns(deps, &altv1.CompareBacktestRunsRequest{}) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(resp.GetResults()) != 0 { t.Errorf("expected empty results, got %d", len(resp.GetResults())) } if resp.GetError() != nil { t.Fatalf("expected no error for empty no-op compare, got %+v", resp.GetError()) } } func TestHandleCompareBacktestRunsRejectsEmptyID(t *testing.T) { deps := BacktestDeps{} resp, err := handleCompareBacktestRuns(deps, &altv1.CompareBacktestRunsRequest{RunIds: []string{"run-1", ""}}) if err != nil { t.Fatalf("expected typed validation response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInvalidRequest) } func TestBacktestHandlersCoverAllRequests(t *testing.T) { deps := BacktestDeps{} registered := make(map[string]bool) for _, h := range backtestHandlers(deps) { if h.requestType == "" { t.Error("handler has empty request type") } if h.register == nil { t.Errorf("handler %q has nil register", h.requestType) } registered[h.requestType] = true } required := []string{ "alt.v1.StartBacktestRequest", "alt.v1.GetBacktestRunRequest", "alt.v1.ListBacktestRunsRequest", "alt.v1.GetBacktestRunDetailRequest", "alt.v1.GetBacktestResultRequest", "alt.v1.CompareBacktestRunsRequest", } for _, req := range required { if !registered[req] { t.Errorf("missing handler for %q", req) } } } func TestWorkerBacktestSocketValidationReturnsTypedError(t *testing.T) { client := startBacktestWorkerTestClient(t, BacktestDeps{Starter: &fakeStarter{}}) resp, err := protoSocket.SendRequestTyped[*altv1.StartBacktestRequest, *altv1.StartBacktestResponse]( &client.Communicator, &altv1.StartBacktestRequest{}, 500*time.Millisecond, ) if err != nil { t.Fatalf("request should return typed error response without timeout: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInvalidRequest) } func TestWorkerBacktestSocketUnavailableReturnsTypedError(t *testing.T) { client := startBacktestWorkerTestClient(t, BacktestDeps{}) resp, err := protoSocket.SendRequestTyped[*altv1.StartBacktestRequest, *altv1.StartBacktestResponse]( &client.Communicator, validStartRequest(), 500*time.Millisecond, ) if err != nil { t.Fatalf("request should return typed error response without timeout: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorUnavailable) } func startBacktestWorkerTestClient(t *testing.T, deps BacktestDeps) *protoSocket.WsClient { t.Helper() listener, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("failed to listen on temporary port: %v", err) } port := listener.Addr().(*net.TCPAddr).Port _ = listener.Close() cfg := config.Config{ Host: "127.0.0.1", Port: port, SocketPath: "/socket", } ctx, cancel := context.WithCancel(context.Background()) t.Cleanup(cancel) server := NewServer(cfg, deps) if err := server.Start(ctx); err != nil { t.Fatalf("failed to start worker socket server: %v", err) } t.Cleanup(func() { _ = server.Stop() }) client, err := protoSocket.DialWsWithHeartbeat(ctx, cfg.Host, cfg.Port, cfg.SocketPath, 0, 0, workerContracts.ParserMap()) if err != nil { t.Fatalf("failed to dial worker socket server: %v", err) } t.Cleanup(func() { _ = client.Close() }) return client } func requireBacktestError(t *testing.T, errInfo *altv1.ErrorInfo, code string) { t.Helper() if errInfo == nil { t.Fatalf("expected ErrorInfo code %q, got nil", code) } if errInfo.GetCode() != code { t.Fatalf("expected ErrorInfo code %q, got %q (%s)", code, errInfo.GetCode(), errInfo.GetMessage()) } if errInfo.GetMessage() == "" { t.Fatalf("expected ErrorInfo message for code %q", code) } }