alt/services/api/internal/socket/backtest_test.go
toki bdbfe8228e feat: scheduled market data refresh - parser_map, refresh status model, CLI output, scheduler backfill, socket handlers
- Add parser_map for market data refresh configuration propagation (cli/worker/api)
- Implement refresh status model with SQLite persistence (status_store.go)
- Add scheduler refresh status headless output to CLI operator
- Extend backfill scheduler with status tracking (start/complete/fail)
- Add socket events for scheduler refresh status (start/complete/fail)
- Update proto definitions for refresh status
- Add generated PB files for client
2026-06-22 18:46:45 +09:00

683 lines
24 KiB
Go

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 (f *fakeWorkerClient) SchedulerRefreshStatus(ctx context.Context, req *altv1.SchedulerRefreshStatusRequest) (*altv1.SchedulerRefreshStatusResponse, error) {
return &altv1.SchedulerRefreshStatusResponse{}, 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")
}
}