- Add backtest proto definitions with selector and state types - Update domain types for backtest session and configuration - Implement backtest selector flow for data range input - Add CLI operator client and handoff test coverage - Update generated code for proto changes - Add worker socket backtest mapping implementations - Include agent-task for milestone tracking
486 lines
16 KiB
Go
486 lines
16 KiB
Go
package workerclient
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"net"
|
|
"testing"
|
|
"time"
|
|
|
|
altv1 "git.toki-labs.com/toki/alt/packages/contracts/gen/go/alt/v1"
|
|
apiContracts "git.toki-labs.com/toki/alt/services/api/internal/contracts"
|
|
protoSocket "git.toki-labs.com/toki/proto-socket/go"
|
|
"nhooyr.io/websocket"
|
|
)
|
|
|
|
func TestWorkerClient_Connect_Unavailable(t *testing.T) {
|
|
// Port that is highly unlikely to have anything listening
|
|
client := New("ws://127.0.0.1:54321/socket")
|
|
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
|
|
defer cancel()
|
|
|
|
err := client.Connect(ctx)
|
|
if err == nil {
|
|
t.Fatalf("expected error on unavailable worker, got nil")
|
|
}
|
|
if !errors.Is(err, ErrUnavailable) {
|
|
t.Errorf("expected ErrUnavailable, got %v", err)
|
|
}
|
|
}
|
|
|
|
func startFakeWorker(t *testing.T, handler func(*protoSocket.WsClient)) (int, func()) {
|
|
l, err := net.Listen("tcp", "127.0.0.1:0")
|
|
if err != nil {
|
|
t.Fatalf("failed to listen on temporary port: %v", err)
|
|
}
|
|
port := l.Addr().(*net.TCPAddr).Port
|
|
l.Close()
|
|
|
|
wsServer := protoSocket.NewWsServer("127.0.0.1", port, "/socket", func(conn *websocket.Conn) *protoSocket.WsClient {
|
|
return protoSocket.NewWsClient(conn, 30, 10, apiContracts.ParserMap())
|
|
})
|
|
wsServer.OnClientConnected = handler
|
|
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
if err := wsServer.Start(ctx); err != nil {
|
|
t.Fatalf("failed to start fake worker server: %v", err)
|
|
}
|
|
|
|
cleanup := func() {
|
|
cancel()
|
|
_ = wsServer.Stop()
|
|
}
|
|
|
|
return port, cleanup
|
|
}
|
|
|
|
func TestWorkerClient_Hello_Success(t *testing.T) {
|
|
port, cleanup := startFakeWorker(t, func(client *protoSocket.WsClient) {
|
|
protoSocket.AddRequestListenerTyped[*altv1.HelloRequest, *altv1.HelloResponse](&client.Communicator, func(req *altv1.HelloRequest) (*altv1.HelloResponse, error) {
|
|
return &altv1.HelloResponse{
|
|
ServerName: "alt-worker-fake",
|
|
ServerVersion: "test",
|
|
AltProtocolVersion: req.GetAltProtocolVersion(),
|
|
}, nil
|
|
})
|
|
})
|
|
defer cleanup()
|
|
|
|
client := New(fmt.Sprintf("ws://127.0.0.1:%d/socket", port))
|
|
ctx := context.Background()
|
|
|
|
if err := client.Connect(ctx); err != nil {
|
|
t.Fatalf("failed to connect: %v", err)
|
|
}
|
|
defer client.Close()
|
|
|
|
res, err := client.Hello(ctx, &altv1.HelloRequest{
|
|
AltProtocolVersion: "alt.v1",
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("Hello request failed: %v", err)
|
|
}
|
|
|
|
if res.ServerName != "alt-worker-fake" {
|
|
t.Errorf("expected ServerName to be alt-worker-fake, got %q", res.ServerName)
|
|
}
|
|
}
|
|
|
|
func TestWorkerClient_StartBacktest_Success(t *testing.T) {
|
|
port, cleanup := startFakeWorker(t, func(client *protoSocket.WsClient) {
|
|
protoSocket.AddRequestListenerTyped[*altv1.StartBacktestRequest, *altv1.StartBacktestResponse](&client.Communicator, func(req *altv1.StartBacktestRequest) (*altv1.StartBacktestResponse, error) {
|
|
return &altv1.StartBacktestResponse{
|
|
Run: &altv1.BacktestRun{
|
|
Id: "run-1",
|
|
Spec: req.GetSpec(),
|
|
Status: altv1.BacktestRunStatus_BACKTEST_RUN_STATUS_PENDING,
|
|
},
|
|
}, nil
|
|
})
|
|
})
|
|
defer cleanup()
|
|
|
|
client := New(fmt.Sprintf("ws://127.0.0.1:%d/socket", port))
|
|
ctx := context.Background()
|
|
if err := client.Connect(ctx); err != nil {
|
|
t.Fatalf("failed to connect: %v", err)
|
|
}
|
|
defer client.Close()
|
|
|
|
res, err := client.StartBacktest(ctx, &altv1.StartBacktestRequest{
|
|
Spec: &altv1.BacktestRunSpec{
|
|
StrategyId: "strat-abc",
|
|
Selector: &altv1.BacktestInputSelector{
|
|
InstrumentIds: []string{"005930"},
|
|
Symbols: []string{"AAPL"},
|
|
},
|
|
},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("StartBacktest failed: %v", err)
|
|
}
|
|
if res.GetRun().GetId() != "run-1" {
|
|
t.Errorf("expected run id run-1, got %q", res.GetRun().GetId())
|
|
}
|
|
if res.GetRun().GetSpec().GetStrategyId() != "strat-abc" {
|
|
t.Errorf("spec did not round-trip, got %q", res.GetRun().GetSpec().GetStrategyId())
|
|
}
|
|
pSel := res.GetRun().GetSpec().GetSelector()
|
|
if pSel == nil {
|
|
t.Fatal("expected selector to round-trip")
|
|
}
|
|
if len(pSel.GetInstrumentIds()) != 1 || pSel.GetInstrumentIds()[0] != "005930" {
|
|
t.Errorf("unexpected selector instrument ids: %v", pSel.GetInstrumentIds())
|
|
}
|
|
if len(pSel.GetSymbols()) != 1 || pSel.GetSymbols()[0] != "AAPL" {
|
|
t.Errorf("unexpected selector symbols: %v", pSel.GetSymbols())
|
|
}
|
|
}
|
|
|
|
func TestWorkerClient_StartPaperTrading_Success(t *testing.T) {
|
|
port, cleanup := startFakeWorker(t, func(client *protoSocket.WsClient) {
|
|
protoSocket.AddRequestListenerTyped[*altv1.StartPaperTradingRequest, *altv1.StartPaperTradingResponse](&client.Communicator, func(req *altv1.StartPaperTradingRequest) (*altv1.StartPaperTradingResponse, error) {
|
|
return &altv1.StartPaperTradingResponse{
|
|
State: &altv1.PaperTradingState{
|
|
AccountId: req.GetAccountId(),
|
|
Cash: req.GetStartingCash(),
|
|
},
|
|
}, nil
|
|
})
|
|
})
|
|
defer cleanup()
|
|
|
|
client := New(fmt.Sprintf("ws://127.0.0.1:%d/socket", port))
|
|
ctx := context.Background()
|
|
if err := client.Connect(ctx); err != nil {
|
|
t.Fatalf("failed to connect: %v", err)
|
|
}
|
|
defer client.Close()
|
|
|
|
res, err := client.StartPaperTrading(ctx, &altv1.StartPaperTradingRequest{
|
|
AccountId: "paper-1",
|
|
Spec: &altv1.BacktestRunSpec{StrategyId: "strategy-v1"},
|
|
StartingCash: &altv1.Price{Currency: altv1.Currency_CURRENCY_KRW, Amount: &altv1.Decimal{Value: "10000000"}},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("StartPaperTrading failed: %v", err)
|
|
}
|
|
if res.GetState().GetAccountId() != "paper-1" {
|
|
t.Errorf("account id did not round-trip, got %q", res.GetState().GetAccountId())
|
|
}
|
|
if res.GetState().GetCash().GetAmount().GetValue() != "10000000" {
|
|
t.Errorf("starting cash did not round-trip, got %q", res.GetState().GetCash().GetAmount().GetValue())
|
|
}
|
|
}
|
|
|
|
func TestWorkerClient_GetPaperTradingState_Success(t *testing.T) {
|
|
port, cleanup := startFakeWorker(t, func(client *protoSocket.WsClient) {
|
|
protoSocket.AddRequestListenerTyped[*altv1.GetPaperTradingStateRequest, *altv1.GetPaperTradingStateResponse](&client.Communicator, func(req *altv1.GetPaperTradingStateRequest) (*altv1.GetPaperTradingStateResponse, error) {
|
|
return &altv1.GetPaperTradingStateResponse{
|
|
State: &altv1.PaperTradingState{AccountId: req.GetAccountId()},
|
|
}, nil
|
|
})
|
|
})
|
|
defer cleanup()
|
|
|
|
client := New(fmt.Sprintf("ws://127.0.0.1:%d/socket", port))
|
|
ctx := context.Background()
|
|
if err := client.Connect(ctx); err != nil {
|
|
t.Fatalf("failed to connect: %v", err)
|
|
}
|
|
defer client.Close()
|
|
|
|
res, err := client.GetPaperTradingState(ctx, &altv1.GetPaperTradingStateRequest{AccountId: "paper-1"})
|
|
if err != nil {
|
|
t.Fatalf("GetPaperTradingState failed: %v", err)
|
|
}
|
|
if res.GetState().GetAccountId() != "paper-1" {
|
|
t.Errorf("account id did not round-trip, got %q", res.GetState().GetAccountId())
|
|
}
|
|
}
|
|
|
|
func TestWorkerClient_GetBacktestRun_Success(t *testing.T) {
|
|
port, cleanup := startFakeWorker(t, func(client *protoSocket.WsClient) {
|
|
protoSocket.AddRequestListenerTyped[*altv1.GetBacktestRunRequest, *altv1.GetBacktestRunResponse](&client.Communicator, func(req *altv1.GetBacktestRunRequest) (*altv1.GetBacktestRunResponse, error) {
|
|
return &altv1.GetBacktestRunResponse{
|
|
Run: &altv1.BacktestRun{
|
|
Id: req.GetRunId(),
|
|
Status: altv1.BacktestRunStatus_BACKTEST_RUN_STATUS_SUCCEEDED,
|
|
},
|
|
}, nil
|
|
})
|
|
})
|
|
defer cleanup()
|
|
|
|
client := New(fmt.Sprintf("ws://127.0.0.1:%d/socket", port))
|
|
ctx := context.Background()
|
|
if err := client.Connect(ctx); err != nil {
|
|
t.Fatalf("failed to connect: %v", err)
|
|
}
|
|
defer client.Close()
|
|
|
|
res, err := client.GetBacktestRun(ctx, &altv1.GetBacktestRunRequest{RunId: "run-1"})
|
|
if err != nil {
|
|
t.Fatalf("GetBacktestRun failed: %v", err)
|
|
}
|
|
if res.GetRun().GetId() != "run-1" {
|
|
t.Errorf("run id did not round-trip, got %q", res.GetRun().GetId())
|
|
}
|
|
if res.GetRun().GetStatus() != altv1.BacktestRunStatus_BACKTEST_RUN_STATUS_SUCCEEDED {
|
|
t.Errorf("status did not round-trip, got %v", res.GetRun().GetStatus())
|
|
}
|
|
}
|
|
|
|
func TestWorkerClient_GetBacktestRun_Unavailable(t *testing.T) {
|
|
// A never-connected client must report ErrUnavailable instead of blocking.
|
|
client := New("ws://127.0.0.1:54321/socket")
|
|
_, err := client.GetBacktestRun(context.Background(), &altv1.GetBacktestRunRequest{RunId: "run-1"})
|
|
if !errors.Is(err, ErrUnavailable) {
|
|
t.Errorf("expected ErrUnavailable, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestWorkerClient_ListInstruments_Success(t *testing.T) {
|
|
port, cleanup := startFakeWorker(t, func(client *protoSocket.WsClient) {
|
|
protoSocket.AddRequestListenerTyped[*altv1.ListInstrumentsRequest, *altv1.ListInstrumentsResponse](&client.Communicator, func(req *altv1.ListInstrumentsRequest) (*altv1.ListInstrumentsResponse, error) {
|
|
return &altv1.ListInstrumentsResponse{
|
|
Instruments: []*altv1.Instrument{{
|
|
Id: "KRX:005930",
|
|
Market: req.GetMarket(),
|
|
Symbol: "005930",
|
|
Currency: altv1.Currency_CURRENCY_KRW,
|
|
}},
|
|
}, nil
|
|
})
|
|
})
|
|
defer cleanup()
|
|
|
|
client := New(fmt.Sprintf("ws://127.0.0.1:%d/socket", port))
|
|
ctx := context.Background()
|
|
if err := client.Connect(ctx); err != nil {
|
|
t.Fatalf("failed to connect: %v", err)
|
|
}
|
|
defer client.Close()
|
|
|
|
res, err := client.ListInstruments(ctx, &altv1.ListInstrumentsRequest{Market: altv1.Market_MARKET_KR})
|
|
if err != nil {
|
|
t.Fatalf("ListInstruments failed: %v", err)
|
|
}
|
|
if len(res.GetInstruments()) != 1 {
|
|
t.Fatalf("expected 1 instrument, got %d", len(res.GetInstruments()))
|
|
}
|
|
if res.GetInstruments()[0].GetMarket() != altv1.Market_MARKET_KR {
|
|
t.Errorf("market did not round-trip, got %v", res.GetInstruments()[0].GetMarket())
|
|
}
|
|
}
|
|
|
|
func TestWorkerClient_ListBars_Success(t *testing.T) {
|
|
port, cleanup := startFakeWorker(t, func(client *protoSocket.WsClient) {
|
|
protoSocket.AddRequestListenerTyped[*altv1.ListBarsRequest, *altv1.ListBarsResponse](&client.Communicator, func(req *altv1.ListBarsRequest) (*altv1.ListBarsResponse, error) {
|
|
return &altv1.ListBarsResponse{
|
|
Bars: []*altv1.Bar{{
|
|
InstrumentId: req.GetInstrumentId(),
|
|
Timeframe: req.GetTimeframe(),
|
|
TimestampUnixMs: req.GetFromUnixMs(),
|
|
}},
|
|
}, nil
|
|
})
|
|
})
|
|
defer cleanup()
|
|
|
|
client := New(fmt.Sprintf("ws://127.0.0.1:%d/socket", port))
|
|
ctx := context.Background()
|
|
if err := client.Connect(ctx); err != nil {
|
|
t.Fatalf("failed to connect: %v", err)
|
|
}
|
|
defer client.Close()
|
|
|
|
res, err := client.ListBars(ctx, &altv1.ListBarsRequest{
|
|
InstrumentId: "KRX:005930",
|
|
Timeframe: altv1.Timeframe_TIMEFRAME_DAILY,
|
|
FromUnixMs: 1,
|
|
ToUnixMs: 2,
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("ListBars failed: %v", err)
|
|
}
|
|
if len(res.GetBars()) != 1 {
|
|
t.Fatalf("expected 1 bar, got %d", len(res.GetBars()))
|
|
}
|
|
if res.GetBars()[0].GetInstrumentId() != "KRX:005930" {
|
|
t.Errorf("instrument id did not round-trip, got %q", res.GetBars()[0].GetInstrumentId())
|
|
}
|
|
}
|
|
|
|
func TestWorkerClient_ImportDailyBars_Success(t *testing.T) {
|
|
port, cleanup := startFakeWorker(t, func(client *protoSocket.WsClient) {
|
|
protoSocket.AddRequestListenerTyped[*altv1.ImportDailyBarsRequest, *altv1.ImportDailyBarsResponse](&client.Communicator, func(req *altv1.ImportDailyBarsRequest) (*altv1.ImportDailyBarsResponse, error) {
|
|
return &altv1.ImportDailyBarsResponse{
|
|
Provider: req.GetProvider(),
|
|
InstrumentCount: int32(len(req.GetSymbols())),
|
|
BarCount: 2,
|
|
}, nil
|
|
})
|
|
})
|
|
defer cleanup()
|
|
|
|
client := New(fmt.Sprintf("ws://127.0.0.1:%d/socket", port))
|
|
ctx := context.Background()
|
|
if err := client.Connect(ctx); err != nil {
|
|
t.Fatalf("failed to connect: %v", err)
|
|
}
|
|
defer client.Close()
|
|
|
|
res, err := client.ImportDailyBars(ctx, &altv1.ImportDailyBarsRequest{
|
|
Provider: "kis",
|
|
SelectorKind: "watchlist",
|
|
Market: altv1.Market_MARKET_KR,
|
|
Venue: altv1.Venue_VENUE_KRX,
|
|
Symbols: []string{"005930"},
|
|
})
|
|
if err != nil {
|
|
t.Fatalf("ImportDailyBars failed: %v", err)
|
|
}
|
|
if res.GetProvider() != "kis" {
|
|
t.Errorf("provider did not round-trip, got %q", res.GetProvider())
|
|
}
|
|
if res.GetInstrumentCount() != 1 || res.GetBarCount() != 2 {
|
|
t.Errorf("counts did not round-trip, got instruments=%d bars=%d", res.GetInstrumentCount(), res.GetBarCount())
|
|
}
|
|
}
|
|
|
|
func TestWorkerClient_ImportDailyBars_Unavailable(t *testing.T) {
|
|
// A never-connected client must report ErrUnavailable instead of blocking.
|
|
client := New("ws://127.0.0.1:54321/socket")
|
|
_, err := client.ImportDailyBars(context.Background(), &altv1.ImportDailyBarsRequest{
|
|
Provider: "kis",
|
|
SelectorKind: "watchlist",
|
|
Symbols: []string{"005930"},
|
|
})
|
|
if !errors.Is(err, ErrUnavailable) {
|
|
t.Errorf("expected ErrUnavailable, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestWorkerClient_StartBacktest_Unavailable(t *testing.T) {
|
|
// A client that was never connected must report ErrUnavailable rather than
|
|
// panicking or blocking, exercising the shared sendTyped guard.
|
|
client := New("ws://127.0.0.1:54321/socket")
|
|
_, err := client.StartBacktest(context.Background(), &altv1.StartBacktestRequest{
|
|
Spec: &altv1.BacktestRunSpec{StrategyId: "strat-abc"},
|
|
})
|
|
if !errors.Is(err, ErrUnavailable) {
|
|
t.Errorf("expected ErrUnavailable, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestWorkerClient_Hello_Timeout(t *testing.T) {
|
|
port, cleanup := startFakeWorker(t, func(client *protoSocket.WsClient) {
|
|
protoSocket.AddRequestListenerTyped[*altv1.HelloRequest, *altv1.HelloResponse](&client.Communicator, func(req *altv1.HelloRequest) (*altv1.HelloResponse, error) {
|
|
// Deliberately delay response to trigger timeout
|
|
time.Sleep(200 * time.Millisecond)
|
|
return &altv1.HelloResponse{
|
|
ServerName: "alt-worker-fake",
|
|
}, nil
|
|
})
|
|
})
|
|
defer cleanup()
|
|
|
|
client := New(fmt.Sprintf("ws://127.0.0.1:%d/socket", port))
|
|
ctx := context.Background()
|
|
|
|
if err := client.Connect(ctx); err != nil {
|
|
t.Fatalf("failed to connect: %v", err)
|
|
}
|
|
defer client.Close()
|
|
|
|
// Short deadline context
|
|
timeoutCtx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond)
|
|
defer cancel()
|
|
|
|
_, err := client.Hello(timeoutCtx, &altv1.HelloRequest{
|
|
AltProtocolVersion: "alt.v1",
|
|
})
|
|
if err == nil {
|
|
t.Fatalf("expected timeout error, got nil")
|
|
}
|
|
if !errors.Is(err, ErrTimeout) {
|
|
t.Errorf("expected ErrTimeout, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestWorkerClient_Hello_ContextCanceled(t *testing.T) {
|
|
port, cleanup := startFakeWorker(t, func(client *protoSocket.WsClient) {})
|
|
defer cleanup()
|
|
|
|
client := New(fmt.Sprintf("ws://127.0.0.1:%d/socket", port))
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
cancel()
|
|
|
|
if err := client.Connect(context.Background()); err != nil {
|
|
t.Fatalf("failed to connect: %v", err)
|
|
}
|
|
defer client.Close()
|
|
|
|
_, err := client.Hello(ctx, &altv1.HelloRequest{AltProtocolVersion: "alt.v1"})
|
|
if err == nil {
|
|
t.Fatalf("expected error on canceled context, got nil")
|
|
}
|
|
if !errors.Is(err, context.Canceled) {
|
|
t.Errorf("expected context.Canceled, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestWorkerClient_Hello_ContextCanceled_Midflight(t *testing.T) {
|
|
port, cleanup := startFakeWorker(t, func(client *protoSocket.WsClient) {
|
|
protoSocket.AddRequestListenerTyped[*altv1.HelloRequest, *altv1.HelloResponse](&client.Communicator, func(req *altv1.HelloRequest) (*altv1.HelloResponse, error) {
|
|
time.Sleep(200 * time.Millisecond)
|
|
return &altv1.HelloResponse{ServerName: "alt-worker-fake"}, nil
|
|
})
|
|
})
|
|
defer cleanup()
|
|
|
|
client := New(fmt.Sprintf("ws://127.0.0.1:%d/socket", port))
|
|
if err := client.Connect(context.Background()); err != nil {
|
|
t.Fatalf("failed to connect: %v", err)
|
|
}
|
|
defer client.Close()
|
|
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
|
|
go func() {
|
|
time.Sleep(50 * time.Millisecond)
|
|
cancel()
|
|
}()
|
|
|
|
_, err := client.Hello(ctx, &altv1.HelloRequest{AltProtocolVersion: "alt.v1"})
|
|
if err == nil {
|
|
t.Fatalf("expected error on canceled context mid-flight, got nil")
|
|
}
|
|
if !errors.Is(err, context.Canceled) {
|
|
t.Errorf("expected context.Canceled, got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestWorkerClient_Hello_DeadlineExceeded(t *testing.T) {
|
|
port, cleanup := startFakeWorker(t, func(client *protoSocket.WsClient) {})
|
|
defer cleanup()
|
|
|
|
client := New(fmt.Sprintf("ws://127.0.0.1:%d/socket", port))
|
|
ctx, cancel := context.WithDeadline(context.Background(), time.Now().Add(-1*time.Second))
|
|
defer cancel()
|
|
|
|
if err := client.Connect(context.Background()); err != nil {
|
|
t.Fatalf("failed to connect: %v", err)
|
|
}
|
|
defer client.Close()
|
|
|
|
_, err := client.Hello(ctx, &altv1.HelloRequest{AltProtocolVersion: "alt.v1"})
|
|
if err == nil {
|
|
t.Fatalf("expected error on exceeded deadline, got nil")
|
|
}
|
|
if !errors.Is(err, ErrTimeout) {
|
|
t.Errorf("expected ErrTimeout, got %v", err)
|
|
}
|
|
}
|