- 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/
376 lines
12 KiB
Go
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)
|
|
}
|
|
}
|