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 listReq *altv1.ListBacktestRunsRequest detailReq *altv1.GetBacktestRunDetailRequest resultReq *altv1.GetBacktestResultRequest compareReq *altv1.CompareBacktestRunsRequest instReq *altv1.ListInstrumentsRequest barsReq *altv1.ListBarsRequest importReq *altv1.ImportDailyBarsRequest startRes *altv1.StartBacktestResponse listRes *altv1.ListBacktestRunsResponse detailRes *altv1.GetBacktestRunDetailResponse resultRes *altv1.GetBacktestResultResponse compareRes *altv1.CompareBacktestRunsResponse instRes *altv1.ListInstrumentsResponse barsRes *altv1.ListBarsResponse importRes *altv1.ImportDailyBarsResponse 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) 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) 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 validStart() *altv1.StartBacktestRequest { return &altv1.StartBacktestRequest{Spec: &altv1.BacktestRunSpec{StrategyId: "strat-abc"}} } 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()) } } 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 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: "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.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.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") } }