package socket import ( "context" "testing" "time" altv1 "git.toki-labs.com/toki/alt/packages/contracts/gen/go/alt/v1" "git.toki-labs.com/toki/alt/services/api/internal/config" apiContracts "git.toki-labs.com/toki/alt/services/api/internal/contracts" "git.toki-labs.com/toki/alt/services/api/internal/workerclient" protoSocket "git.toki-labs.com/toki/proto-socket/go" ) // fakeWorkerClient is a controllable WorkerClient for handler tests. It records // the last forwarded request and returns canned responses/errors. type fakeWorkerClient struct { startReq *altv1.StartBacktestRequest getRunReq *altv1.GetBacktestRunRequest listReq *altv1.ListBacktestRunsRequest detailReq *altv1.GetBacktestRunDetailRequest resultReq *altv1.GetBacktestResultRequest compareReq *altv1.CompareBacktestRunsRequest monthlyReq *altv1.AggregateMonthlyBarsRequest instReq *altv1.ListInstrumentsRequest barsReq *altv1.ListBarsRequest importReq *altv1.ImportDailyBarsRequest paperStartReq *altv1.StartPaperTradingRequest paperStateReq *altv1.GetPaperTradingStateRequest paperSubmitReq *altv1.SubmitPaperOrderRequest paperCancelReq *altv1.CancelPaperOrderRequest paperFillReq *altv1.FillPaperOrderRequest liveCapReq *altv1.GetLiveBrokerCapabilitiesRequest liveSubmitReq *altv1.SubmitLiveOrderRequest liveCancelReq *altv1.CancelLiveOrderRequest liveGetReq *altv1.GetLiveOrderRequest liveRiskPolicyReq *altv1.GetLiveRiskPolicyRequest liveGetKillSwitchReq *altv1.GetLiveKillSwitchRequest liveSetKillSwitchReq *altv1.SetLiveKillSwitchRequest liveSyncAccReq *altv1.SyncLiveAccountRequest liveGetAccSnapReq *altv1.GetLiveAccountSnapshotRequest liveListAuditReq *altv1.ListLiveAuditEventsRequest startRes *altv1.StartBacktestResponse getRunRes *altv1.GetBacktestRunResponse listRes *altv1.ListBacktestRunsResponse detailRes *altv1.GetBacktestRunDetailResponse resultRes *altv1.GetBacktestResultResponse compareRes *altv1.CompareBacktestRunsResponse monthlyRes *altv1.AggregateMonthlyBarsResponse instRes *altv1.ListInstrumentsResponse barsRes *altv1.ListBarsResponse importRes *altv1.ImportDailyBarsResponse paperStartRes *altv1.StartPaperTradingResponse paperStateRes *altv1.GetPaperTradingStateResponse paperSubmitRes *altv1.SubmitPaperOrderResponse paperCancelRes *altv1.CancelPaperOrderResponse paperFillRes *altv1.FillPaperOrderResponse liveCapRes *altv1.GetLiveBrokerCapabilitiesResponse liveSubmitRes *altv1.SubmitLiveOrderResponse liveCancelRes *altv1.CancelLiveOrderResponse liveGetRes *altv1.GetLiveOrderResponse liveRiskPolicyRes *altv1.GetLiveRiskPolicyResponse liveGetKillSwitchRes *altv1.GetLiveKillSwitchResponse liveSetKillSwitchRes *altv1.SetLiveKillSwitchResponse liveSyncAccRes *altv1.SyncLiveAccountResponse liveGetAccSnapRes *altv1.GetLiveAccountSnapshotResponse liveListAuditRes *altv1.ListLiveAuditEventsResponse err error connectErr error isConnected bool connectCount int } func (f *fakeWorkerClient) IsConnected() bool { return f.isConnected } func (f *fakeWorkerClient) Connect(ctx context.Context) error { f.connectCount++ if f.connectErr != nil { return f.connectErr } f.isConnected = true return nil } func (f *fakeWorkerClient) Close() error { f.isConnected = false return nil } func (f *fakeWorkerClient) Hello(ctx context.Context, req *altv1.HelloRequest) (*altv1.HelloResponse, error) { return &altv1.HelloResponse{}, f.err } func (f *fakeWorkerClient) StartBacktest(ctx context.Context, req *altv1.StartBacktestRequest) (*altv1.StartBacktestResponse, error) { f.startReq = req return f.startRes, f.err } func (f *fakeWorkerClient) GetBacktestRun(ctx context.Context, req *altv1.GetBacktestRunRequest) (*altv1.GetBacktestRunResponse, error) { f.getRunReq = req return f.getRunRes, f.err } func (f *fakeWorkerClient) ListBacktestRuns(ctx context.Context, req *altv1.ListBacktestRunsRequest) (*altv1.ListBacktestRunsResponse, error) { f.listReq = req return f.listRes, f.err } func (f *fakeWorkerClient) GetBacktestRunDetail(ctx context.Context, req *altv1.GetBacktestRunDetailRequest) (*altv1.GetBacktestRunDetailResponse, error) { f.detailReq = req return f.detailRes, f.err } func (f *fakeWorkerClient) GetBacktestResult(ctx context.Context, req *altv1.GetBacktestResultRequest) (*altv1.GetBacktestResultResponse, error) { f.resultReq = req return f.resultRes, f.err } func (f *fakeWorkerClient) CompareBacktestRuns(ctx context.Context, req *altv1.CompareBacktestRunsRequest) (*altv1.CompareBacktestRunsResponse, error) { f.compareReq = req return f.compareRes, f.err } func (f *fakeWorkerClient) StartPaperTrading(ctx context.Context, req *altv1.StartPaperTradingRequest) (*altv1.StartPaperTradingResponse, error) { f.paperStartReq = req return f.paperStartRes, f.err } func (f *fakeWorkerClient) GetPaperTradingState(ctx context.Context, req *altv1.GetPaperTradingStateRequest) (*altv1.GetPaperTradingStateResponse, error) { f.paperStateReq = req return f.paperStateRes, f.err } func (f *fakeWorkerClient) SubmitPaperOrder(ctx context.Context, req *altv1.SubmitPaperOrderRequest) (*altv1.SubmitPaperOrderResponse, error) { f.paperSubmitReq = req return f.paperSubmitRes, f.err } func (f *fakeWorkerClient) CancelPaperOrder(ctx context.Context, req *altv1.CancelPaperOrderRequest) (*altv1.CancelPaperOrderResponse, error) { f.paperCancelReq = req return f.paperCancelRes, f.err } func (f *fakeWorkerClient) FillPaperOrder(ctx context.Context, req *altv1.FillPaperOrderRequest) (*altv1.FillPaperOrderResponse, error) { f.paperFillReq = req return f.paperFillRes, f.err } func (f *fakeWorkerClient) ListInstruments(ctx context.Context, req *altv1.ListInstrumentsRequest) (*altv1.ListInstrumentsResponse, error) { f.instReq = req return f.instRes, f.err } func (f *fakeWorkerClient) ListBars(ctx context.Context, req *altv1.ListBarsRequest) (*altv1.ListBarsResponse, error) { f.barsReq = req return f.barsRes, f.err } func (f *fakeWorkerClient) ImportDailyBars(ctx context.Context, req *altv1.ImportDailyBarsRequest) (*altv1.ImportDailyBarsResponse, error) { f.importReq = req return f.importRes, f.err } func (f *fakeWorkerClient) AggregateMonthlyBars(ctx context.Context, req *altv1.AggregateMonthlyBarsRequest) (*altv1.AggregateMonthlyBarsResponse, error) { f.monthlyReq = req return f.monthlyRes, f.err } func (f *fakeWorkerClient) GetLiveBrokerCapabilities(ctx context.Context, req *altv1.GetLiveBrokerCapabilitiesRequest) (*altv1.GetLiveBrokerCapabilitiesResponse, error) { f.liveCapReq = req return f.liveCapRes, f.err } func (f *fakeWorkerClient) GetLiveBrokerCapabilitiesCallCount() int { if f.liveCapReq == nil { return 0 } return 1 } func (f *fakeWorkerClient) SubmitLiveOrder(ctx context.Context, req *altv1.SubmitLiveOrderRequest) (*altv1.SubmitLiveOrderResponse, error) { f.liveSubmitReq = req return f.liveSubmitRes, f.err } func (f *fakeWorkerClient) CancelLiveOrder(ctx context.Context, req *altv1.CancelLiveOrderRequest) (*altv1.CancelLiveOrderResponse, error) { f.liveCancelReq = req return f.liveCancelRes, f.err } func (f *fakeWorkerClient) GetLiveOrder(ctx context.Context, req *altv1.GetLiveOrderRequest) (*altv1.GetLiveOrderResponse, error) { f.liveGetReq = req return f.liveGetRes, f.err } func (f *fakeWorkerClient) GetLiveRiskPolicy(ctx context.Context, req *altv1.GetLiveRiskPolicyRequest) (*altv1.GetLiveRiskPolicyResponse, error) { f.liveRiskPolicyReq = req return f.liveRiskPolicyRes, f.err } func (f *fakeWorkerClient) GetLiveKillSwitch(ctx context.Context, req *altv1.GetLiveKillSwitchRequest) (*altv1.GetLiveKillSwitchResponse, error) { f.liveGetKillSwitchReq = req return f.liveGetKillSwitchRes, f.err } func (f *fakeWorkerClient) SetLiveKillSwitch(ctx context.Context, req *altv1.SetLiveKillSwitchRequest) (*altv1.SetLiveKillSwitchResponse, error) { f.liveSetKillSwitchReq = req return f.liveSetKillSwitchRes, f.err } func (f *fakeWorkerClient) SyncLiveAccount(ctx context.Context, req *altv1.SyncLiveAccountRequest) (*altv1.SyncLiveAccountResponse, error) { f.liveSyncAccReq = req return f.liveSyncAccRes, f.err } func (f *fakeWorkerClient) GetLiveAccountSnapshot(ctx context.Context, req *altv1.GetLiveAccountSnapshotRequest) (*altv1.GetLiveAccountSnapshotResponse, error) { f.liveGetAccSnapReq = req return f.liveGetAccSnapRes, f.err } func (f *fakeWorkerClient) ListLiveAuditEvents(ctx context.Context, req *altv1.ListLiveAuditEventsRequest) (*altv1.ListLiveAuditEventsResponse, error) { f.liveListAuditReq = req return f.liveListAuditRes, f.err } func validStart() *altv1.StartBacktestRequest { return &altv1.StartBacktestRequest{ Spec: &altv1.BacktestRunSpec{ StrategyId: "strat-abc", Selector: &altv1.BacktestInputSelector{ InstrumentIds: []string{"005930"}, Symbols: []string{"AAPL"}, }, }, } } func TestHandleStartBacktestForwards(t *testing.T) { fake := &fakeWorkerClient{startRes: &altv1.StartBacktestResponse{Run: &altv1.BacktestRun{Id: "run-1"}}} resp, err := handleStartBacktest(fake, validStart()) if err != nil { t.Fatalf("unexpected error: %v", err) } if fake.startReq == nil { t.Fatal("expected request to be forwarded to worker") } if resp.GetRun().GetId() != "run-1" { t.Errorf("unexpected response run id: %q", resp.GetRun().GetId()) } pSel := fake.startReq.GetSpec().GetSelector() if pSel == nil { t.Fatal("expected selector to be forwarded") } if len(pSel.GetInstrumentIds()) != 1 || pSel.GetInstrumentIds()[0] != "005930" { t.Errorf("unexpected forwarded instrument ids: %v", pSel.GetInstrumentIds()) } if len(pSel.GetSymbols()) != 1 || pSel.GetSymbols()[0] != "AAPL" { t.Errorf("unexpected forwarded symbols: %v", pSel.GetSymbols()) } } func TestHandleStartBacktestRejectsMissingSpec(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleStartBacktest(fake, &altv1.StartBacktestRequest{}) if err != nil { t.Fatalf("expected typed validation response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInvalidRequest) if fake.startReq != nil { t.Error("worker must not be called when validation fails") } } func TestHandleStartBacktestNilWorker(t *testing.T) { resp, err := handleStartBacktest(nil, validStart()) if err != nil { t.Fatalf("expected typed unavailable response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorUnavailable) } func TestHandleStartBacktestMapsUnavailable(t *testing.T) { fake := &fakeWorkerClient{err: workerclient.ErrUnavailable} resp, err := handleStartBacktest(fake, validStart()) if err != nil { t.Fatalf("expected typed unavailable response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorUnavailable) } func TestHandleStartBacktestMapsTimeout(t *testing.T) { fake := &fakeWorkerClient{err: workerclient.ErrTimeout} resp, err := handleStartBacktest(fake, validStart()) if err != nil { t.Fatalf("expected typed timeout response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorTimeout) } func TestHandleStartBacktestMapsUnexpectedWorkerError(t *testing.T) { fake := &fakeWorkerClient{err: context.Canceled} resp, err := handleStartBacktest(fake, validStart()) if err != nil { t.Fatalf("expected typed internal response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInternal) } func TestHandleStartBacktestNilWorkerResponseForwardsRequest(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleStartBacktest(fake, validStart()) if err != nil { t.Fatalf("expected typed internal response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInternal) if fake.startReq == nil { t.Error("expected worker to be called before nil response was detected") } } func TestHandleListBacktestRunsNilWorkerResponse(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleListBacktestRuns(fake, &altv1.ListBacktestRunsRequest{}) if err != nil { t.Fatalf("expected typed internal response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInternal) if fake.listReq == nil { t.Error("expected request to be forwarded before nil response was detected") } } func TestHandleListBacktestRunsForwards(t *testing.T) { fake := &fakeWorkerClient{listRes: &altv1.ListBacktestRunsResponse{Runs: []*altv1.BacktestRun{{Id: "run-1"}}}} resp, err := handleListBacktestRuns(fake, &altv1.ListBacktestRunsRequest{Status: altv1.BacktestRunStatus_BACKTEST_RUN_STATUS_SUCCEEDED}) if err != nil { t.Fatalf("unexpected error: %v", err) } if fake.listReq.GetStatus() != altv1.BacktestRunStatus_BACKTEST_RUN_STATUS_SUCCEEDED { t.Error("status filter was not forwarded unchanged") } if len(resp.GetRuns()) != 1 { t.Errorf("expected 1 run, got %d", len(resp.GetRuns())) } } func TestHandleGetBacktestRunForwards(t *testing.T) { fake := &fakeWorkerClient{getRunRes: &altv1.GetBacktestRunResponse{Run: &altv1.BacktestRun{Id: "run-1", Status: altv1.BacktestRunStatus_BACKTEST_RUN_STATUS_RUNNING}}} resp, err := handleGetBacktestRun(fake, &altv1.GetBacktestRunRequest{RunId: "run-1"}) if err != nil { t.Fatalf("unexpected error: %v", err) } if fake.getRunReq.GetRunId() != "run-1" { t.Errorf("run id was not forwarded unchanged, got %q", fake.getRunReq.GetRunId()) } if resp.GetRun().GetId() != "run-1" { t.Errorf("unexpected run id: %q", resp.GetRun().GetId()) } } func TestHandleGetBacktestRunRequiresRunID(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleGetBacktestRun(fake, &altv1.GetBacktestRunRequest{}) if err != nil { t.Fatalf("expected typed validation response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInvalidRequest) if fake.getRunReq != nil { t.Error("worker must not be called when validation fails") } } func TestHandleGetBacktestRunNilWorkerResponse(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleGetBacktestRun(fake, &altv1.GetBacktestRunRequest{RunId: "run-1"}) if err != nil { t.Fatalf("expected typed internal response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInternal) if fake.getRunReq == nil { t.Error("expected request to be forwarded before nil response was detected") } } func TestHandleGetBacktestRunDetailRequiresRunID(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleGetBacktestRunDetail(fake, &altv1.GetBacktestRunDetailRequest{}) if err != nil { t.Fatalf("expected typed validation response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInvalidRequest) if fake.detailReq != nil { t.Error("worker must not be called when validation fails") } } func TestHandleGetBacktestRunDetailNilWorkerResponse(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleGetBacktestRunDetail(fake, &altv1.GetBacktestRunDetailRequest{RunId: "run-1"}) if err != nil { t.Fatalf("expected typed internal response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInternal) if fake.detailReq == nil { t.Error("expected request to be forwarded before nil response was detected") } } func TestHandleGetBacktestRunDetailForwards(t *testing.T) { fake := &fakeWorkerClient{detailRes: &altv1.GetBacktestRunDetailResponse{Run: &altv1.BacktestRun{Id: "run-1"}}} resp, err := handleGetBacktestRunDetail(fake, &altv1.GetBacktestRunDetailRequest{RunId: "run-1"}) if err != nil { t.Fatalf("unexpected error: %v", err) } if resp.GetRun().GetId() != "run-1" { t.Errorf("unexpected run id: %q", resp.GetRun().GetId()) } } func TestHandleGetBacktestResultRequiresRunID(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleGetBacktestResult(fake, &altv1.GetBacktestResultRequest{}) if err != nil { t.Fatalf("expected typed validation response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInvalidRequest) if fake.resultReq != nil { t.Error("worker must not be called when validation fails") } } func TestHandleCompareBacktestRunsEmptyIsNoOp(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleCompareBacktestRuns(fake, &altv1.CompareBacktestRunsRequest{}) if err != nil { t.Fatalf("unexpected error: %v", err) } if resp.GetError() != nil { t.Fatalf("expected no error for empty no-op compare, got %+v", resp.GetError()) } if len(resp.GetResults()) != 0 { t.Errorf("expected empty results, got %d", len(resp.GetResults())) } if fake.compareReq != nil { t.Error("worker must not be called for empty no-op compare") } } func TestHandleCompareBacktestRunsValidation(t *testing.T) { fake := &fakeWorkerClient{} resp, err := handleCompareBacktestRuns(fake, &altv1.CompareBacktestRunsRequest{RunIds: []string{"run-1", ""}}) if err != nil { t.Fatalf("expected typed validation response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorInvalidRequest) if fake.compareReq != nil { t.Error("worker must not be called when validation fails") } } func TestHandleCompareBacktestRunsForwards(t *testing.T) { fake := &fakeWorkerClient{compareRes: &altv1.CompareBacktestRunsResponse{Results: []*altv1.BacktestResult{{RunId: "run-1"}, {RunId: "run-2"}}}} resp, err := handleCompareBacktestRuns(fake, &altv1.CompareBacktestRunsRequest{RunIds: []string{"run-1", "run-2"}}) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(resp.GetResults()) != 2 { t.Errorf("expected 2 results, got %d", len(resp.GetResults())) } } func TestBacktestSocketValidationReturnsTypedErrors(t *testing.T) { fake := &fakeWorkerClient{} client := startBacktestAPITestClient(t, fake) cases := []struct { name string run func(*testing.T) *altv1.ErrorInfo }{ { name: "start missing spec", run: func(t *testing.T) *altv1.ErrorInfo { t.Helper() 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) } return resp.GetError() }, }, { name: "get run missing run id", run: func(t *testing.T) *altv1.ErrorInfo { t.Helper() resp, err := protoSocket.SendRequestTyped[*altv1.GetBacktestRunRequest, *altv1.GetBacktestRunResponse]( &client.Communicator, &altv1.GetBacktestRunRequest{}, 500*time.Millisecond, ) if err != nil { t.Fatalf("request should return typed error response without timeout: %v", err) } return resp.GetError() }, }, { name: "detail missing run id", run: func(t *testing.T) *altv1.ErrorInfo { t.Helper() resp, err := protoSocket.SendRequestTyped[*altv1.GetBacktestRunDetailRequest, *altv1.GetBacktestRunDetailResponse]( &client.Communicator, &altv1.GetBacktestRunDetailRequest{}, 500*time.Millisecond, ) if err != nil { t.Fatalf("request should return typed error response without timeout: %v", err) } return resp.GetError() }, }, { name: "result missing run id", run: func(t *testing.T) *altv1.ErrorInfo { t.Helper() resp, err := protoSocket.SendRequestTyped[*altv1.GetBacktestResultRequest, *altv1.GetBacktestResultResponse]( &client.Communicator, &altv1.GetBacktestResultRequest{}, 500*time.Millisecond, ) if err != nil { t.Fatalf("request should return typed error response without timeout: %v", err) } return resp.GetError() }, }, { name: "compare empty id", run: func(t *testing.T) *altv1.ErrorInfo { t.Helper() resp, err := protoSocket.SendRequestTyped[*altv1.CompareBacktestRunsRequest, *altv1.CompareBacktestRunsResponse]( &client.Communicator, &altv1.CompareBacktestRunsRequest{RunIds: []string{"run-1", ""}}, 500*time.Millisecond, ) if err != nil { t.Fatalf("request should return typed error response without timeout: %v", err) } return resp.GetError() }, }, } for _, tt := range cases { t.Run(tt.name, func(t *testing.T) { requireBacktestError(t, tt.run(t), backtestErrorInvalidRequest) }) } if fake.startReq != nil || fake.getRunReq != nil || fake.detailReq != nil || fake.resultReq != nil || fake.compareReq != nil { t.Fatal("worker must not be called for socket-level validation failures") } } func TestBacktestSocketWorkerUnavailableReturnsTypedError(t *testing.T) { fake := &fakeWorkerClient{err: workerclient.ErrUnavailable} client := startBacktestAPITestClient(t, fake) resp, err := protoSocket.SendRequestTyped[*altv1.StartBacktestRequest, *altv1.StartBacktestResponse]( &client.Communicator, validStart(), 500*time.Millisecond, ) if err != nil { t.Fatalf("request should return typed error response without timeout: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorUnavailable) if fake.startReq == nil { t.Fatal("valid start request should be forwarded before worker error is mapped") } } func startBacktestAPITestClient(t *testing.T, fake workerclient.WorkerClient) *protoSocket.WsClient { t.Helper() ctx, cancel := context.WithCancel(context.Background()) t.Cleanup(cancel) cfg := config.Config{ Host: "127.0.0.1", Port: freeTCPPort(t), SocketPath: "/socket", HeartbeatIntervalSec: 0, HeartbeatWaitSec: 0, } server := NewServerWithWorker(cfg, fake) if err := server.Start(ctx); err != nil { t.Fatalf("failed to start API socket server: %v", err) } t.Cleanup(func() { _ = server.Stop() }) client, err := protoSocket.DialWsWithHeartbeat(ctx, cfg.Host, cfg.Port, cfg.SocketPath, 0, 0, apiContracts.ParserMap()) if err != nil { t.Fatalf("failed to dial API 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) } } func TestBacktestHandlersRegisteredInSession(t *testing.T) { registered := make(map[string]bool) for _, h := range sessionHandlers(nil) { registered[h.requestType] = true } required := []string{ protoSocket.TypeNameOf(&altv1.StartBacktestRequest{}), protoSocket.TypeNameOf(&altv1.GetBacktestRunRequest{}), protoSocket.TypeNameOf(&altv1.ListBacktestRunsRequest{}), protoSocket.TypeNameOf(&altv1.GetBacktestRunDetailRequest{}), protoSocket.TypeNameOf(&altv1.GetBacktestResultRequest{}), protoSocket.TypeNameOf(&altv1.CompareBacktestRunsRequest{}), } for _, req := range required { if !registered[req] { t.Errorf("missing API handler registration for %q", req) } } } func TestHandleStartBacktest_ConnectFailure(t *testing.T) { fake := &fakeWorkerClient{ connectErr: workerclient.ErrUnavailable, isConnected: false, } resp, err := handleStartBacktest(fake, validStart()) if err != nil { t.Fatalf("expected typed unavailable response, got error: %v", err) } requireBacktestError(t, resp.GetError(), backtestErrorUnavailable) if fake.startReq != nil { t.Error("expected request not to be forwarded to worker on connect failure") } if fake.connectCount != 1 { t.Errorf("expected connectCount to be 1, got %d", fake.connectCount) } } func TestHandleStartBacktest_ReconnectBehavior(t *testing.T) { fake := &fakeWorkerClient{ startRes: &altv1.StartBacktestResponse{Run: &altv1.BacktestRun{Id: "run-1"}}, isConnected: false, } resp, err := handleStartBacktest(fake, validStart()) if err != nil { t.Fatalf("unexpected error: %v", err) } if resp.GetError() != nil { t.Fatalf("unexpected handler error: %+v", resp.GetError()) } if fake.connectCount != 1 { t.Errorf("expected Connect to be called once, got %d", fake.connectCount) } if fake.startReq == nil { t.Error("expected request to be forwarded to worker") } if !fake.isConnected { t.Error("expected fake worker state to be connected") } }