Paper trading readiness에서 API/worker/CLI가 같은 protobuf 계약으로 paper state를 시작하고 조회할 수 있어야 한다. Headless 운영 경로를 먼저 닫기 위해 contract, worker runtime, API forwarding, CLI scenario, client parser map과 검증 artifact를 함께 반영한다.
470 lines
16 KiB
Go
470 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"},
|
|
})
|
|
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())
|
|
}
|
|
}
|
|
|
|
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)
|
|
}
|
|
}
|