alt/services/worker/internal/socket/backtest_test.go
toki ec938475e0 chore: socket rail implementation and task archival
- Add backtest and market socket handlers for API and Worker
- Add proto definitions for backtest and market services
- Update client socket integration and tests
- Generate protobuf code for Go and Dart
- Archive completed agent tasks under agent-task/archive/2026/05/
2026-05-30 22:57:18 +09:00

376 lines
12 KiB
Go

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 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.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)
}
}