Paper trading readiness에서 API/worker/CLI가 같은 protobuf 계약으로 paper state를 시작하고 조회할 수 있어야 한다. Headless 운영 경로를 먼저 닫기 위해 contract, worker runtime, API forwarding, CLI scenario, client parser map과 검증 artifact를 함께 반영한다.
449 lines
15 KiB
Go
449 lines
15 KiB
Go
package operator
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"net"
|
|
"sync"
|
|
"testing"
|
|
|
|
"nhooyr.io/websocket"
|
|
|
|
altv1 "git.toki-labs.com/toki/alt/packages/contracts/gen/go/alt/v1"
|
|
protoSocket "git.toki-labs.com/toki/proto-socket/go"
|
|
)
|
|
|
|
// fakeAPI configures the canned responses an in-process API server returns and
|
|
// captures the requests it received so tests can assert the client sent the
|
|
// right fields. A nil response slot makes the handler return an empty typed
|
|
// response.
|
|
type fakeAPI struct {
|
|
helloCaps []string
|
|
instResp *altv1.ListInstrumentsResponse
|
|
barsResp *altv1.ListBarsResponse
|
|
|
|
startBacktestResp *altv1.StartBacktestResponse
|
|
getBacktestRunResp *altv1.GetBacktestRunResponse
|
|
listBacktestRunsResp *altv1.ListBacktestRunsResponse
|
|
getBacktestRunDetailResp *altv1.GetBacktestRunDetailResponse
|
|
getBacktestResultResp *altv1.GetBacktestResultResponse
|
|
compareBacktestRunsResp *altv1.CompareBacktestRunsResponse
|
|
importDailyBarsResp *altv1.ImportDailyBarsResponse
|
|
startPaperResp *altv1.StartPaperTradingResponse
|
|
getPaperStateResp *altv1.GetPaperTradingStateResponse
|
|
|
|
getBacktestRunFunc func(req *altv1.GetBacktestRunRequest) (*altv1.GetBacktestRunResponse, error)
|
|
|
|
mu sync.Mutex
|
|
instReq *altv1.ListInstrumentsRequest
|
|
barsReq *altv1.ListBarsRequest
|
|
startBacktestReq *altv1.StartBacktestRequest
|
|
getBacktestRunReq *altv1.GetBacktestRunRequest
|
|
importDailyBarsReq *altv1.ImportDailyBarsRequest
|
|
startPaperReq *altv1.StartPaperTradingRequest
|
|
getPaperStateReq *altv1.GetPaperTradingStateRequest
|
|
calls []string
|
|
}
|
|
|
|
// recordCall appends an action name in receive order so tests can assert that a
|
|
// command-first scenario issued its import before any status read.
|
|
func (f *fakeAPI) recordCall(action string) {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
f.calls = append(f.calls, action)
|
|
}
|
|
|
|
func (f *fakeAPI) callOrder() []string {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
return append([]string(nil), f.calls...)
|
|
}
|
|
|
|
func (f *fakeAPI) setInstReq(req *altv1.ListInstrumentsRequest) {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
f.instReq = req
|
|
}
|
|
|
|
func (f *fakeAPI) setBarsReq(req *altv1.ListBarsRequest) {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
f.barsReq = req
|
|
}
|
|
|
|
func (f *fakeAPI) setStartBacktestReq(req *altv1.StartBacktestRequest) {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
f.startBacktestReq = req
|
|
}
|
|
|
|
func (f *fakeAPI) setGetBacktestRunReq(req *altv1.GetBacktestRunRequest) {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
f.getBacktestRunReq = req
|
|
}
|
|
|
|
func (f *fakeAPI) lastInstReq() *altv1.ListInstrumentsRequest {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
return f.instReq
|
|
}
|
|
|
|
func (f *fakeAPI) lastBarsReq() *altv1.ListBarsRequest {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
return f.barsReq
|
|
}
|
|
|
|
func (f *fakeAPI) lastStartBacktestReq() *altv1.StartBacktestRequest {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
return f.startBacktestReq
|
|
}
|
|
|
|
func (f *fakeAPI) lastGetBacktestRunReq() *altv1.GetBacktestRunRequest {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
return f.getBacktestRunReq
|
|
}
|
|
|
|
func (f *fakeAPI) setImportDailyBarsReq(req *altv1.ImportDailyBarsRequest) {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
f.importDailyBarsReq = req
|
|
}
|
|
|
|
func (f *fakeAPI) lastImportDailyBarsReq() *altv1.ImportDailyBarsRequest {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
return f.importDailyBarsReq
|
|
}
|
|
|
|
func (f *fakeAPI) setStartPaperReq(req *altv1.StartPaperTradingRequest) {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
f.startPaperReq = req
|
|
}
|
|
|
|
func (f *fakeAPI) lastStartPaperReq() *altv1.StartPaperTradingRequest {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
return f.startPaperReq
|
|
}
|
|
|
|
func (f *fakeAPI) setGetPaperStateReq(req *altv1.GetPaperTradingStateRequest) {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
f.getPaperStateReq = req
|
|
}
|
|
|
|
func (f *fakeAPI) lastGetPaperStateReq() *altv1.GetPaperTradingStateRequest {
|
|
f.mu.Lock()
|
|
defer f.mu.Unlock()
|
|
return f.getPaperStateReq
|
|
}
|
|
|
|
// startFakeAPIServer starts a proto-socket server on a free localhost port that
|
|
// answers hello/list_instruments/list_bars with the supplied canned responses
|
|
// and records the received requests. It returns the ws:// URL and stops the
|
|
// server on test cleanup.
|
|
func startFakeAPIServer(t *testing.T, api *fakeAPI) string {
|
|
t.Helper()
|
|
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
t.Cleanup(cancel)
|
|
|
|
host := "127.0.0.1"
|
|
port := freePort(t)
|
|
path := "/socket"
|
|
|
|
srv := protoSocket.NewWsServer(host, port, path, func(conn *websocket.Conn) *protoSocket.WsClient {
|
|
return protoSocket.NewWsClient(conn, 0, 0, ParserMap())
|
|
})
|
|
srv.OnClientConnected = func(c *protoSocket.WsClient) {
|
|
protoSocket.AddRequestListenerTyped[*altv1.HelloRequest, *altv1.HelloResponse](&c.Communicator, func(req *altv1.HelloRequest) (*altv1.HelloResponse, error) {
|
|
return &altv1.HelloResponse{
|
|
ServerName: "alt-api-fake",
|
|
ServerVersion: "test",
|
|
AltProtocolVersion: req.GetAltProtocolVersion(),
|
|
Capabilities: api.helloCaps,
|
|
}, nil
|
|
})
|
|
protoSocket.AddRequestListenerTyped[*altv1.ListInstrumentsRequest, *altv1.ListInstrumentsResponse](&c.Communicator, func(req *altv1.ListInstrumentsRequest) (*altv1.ListInstrumentsResponse, error) {
|
|
api.recordCall("list_instruments")
|
|
api.setInstReq(req)
|
|
if api.instResp != nil {
|
|
return api.instResp, nil
|
|
}
|
|
return &altv1.ListInstrumentsResponse{}, nil
|
|
})
|
|
protoSocket.AddRequestListenerTyped[*altv1.ListBarsRequest, *altv1.ListBarsResponse](&c.Communicator, func(req *altv1.ListBarsRequest) (*altv1.ListBarsResponse, error) {
|
|
api.recordCall("list_bars")
|
|
api.setBarsReq(req)
|
|
if api.barsResp != nil {
|
|
return api.barsResp, nil
|
|
}
|
|
return &altv1.ListBarsResponse{}, nil
|
|
})
|
|
protoSocket.AddRequestListenerTyped[*altv1.StartBacktestRequest, *altv1.StartBacktestResponse](&c.Communicator, func(req *altv1.StartBacktestRequest) (*altv1.StartBacktestResponse, error) {
|
|
api.setStartBacktestReq(req)
|
|
if api.startBacktestResp != nil {
|
|
return api.startBacktestResp, nil
|
|
}
|
|
return &altv1.StartBacktestResponse{}, nil
|
|
})
|
|
protoSocket.AddRequestListenerTyped[*altv1.GetBacktestRunRequest, *altv1.GetBacktestRunResponse](&c.Communicator, func(req *altv1.GetBacktestRunRequest) (*altv1.GetBacktestRunResponse, error) {
|
|
api.setGetBacktestRunReq(req)
|
|
if api.getBacktestRunFunc != nil {
|
|
return api.getBacktestRunFunc(req)
|
|
}
|
|
if api.getBacktestRunResp != nil {
|
|
return api.getBacktestRunResp, nil
|
|
}
|
|
return &altv1.GetBacktestRunResponse{}, nil
|
|
})
|
|
protoSocket.AddRequestListenerTyped[*altv1.ListBacktestRunsRequest, *altv1.ListBacktestRunsResponse](&c.Communicator, func(req *altv1.ListBacktestRunsRequest) (*altv1.ListBacktestRunsResponse, error) {
|
|
if api.listBacktestRunsResp != nil {
|
|
return api.listBacktestRunsResp, nil
|
|
}
|
|
return &altv1.ListBacktestRunsResponse{}, nil
|
|
})
|
|
protoSocket.AddRequestListenerTyped[*altv1.GetBacktestRunDetailRequest, *altv1.GetBacktestRunDetailResponse](&c.Communicator, func(req *altv1.GetBacktestRunDetailRequest) (*altv1.GetBacktestRunDetailResponse, error) {
|
|
if api.getBacktestRunDetailResp != nil {
|
|
return api.getBacktestRunDetailResp, nil
|
|
}
|
|
return &altv1.GetBacktestRunDetailResponse{}, nil
|
|
})
|
|
protoSocket.AddRequestListenerTyped[*altv1.GetBacktestResultRequest, *altv1.GetBacktestResultResponse](&c.Communicator, func(req *altv1.GetBacktestResultRequest) (*altv1.GetBacktestResultResponse, error) {
|
|
if api.getBacktestResultResp != nil {
|
|
return api.getBacktestResultResp, nil
|
|
}
|
|
return &altv1.GetBacktestResultResponse{}, nil
|
|
})
|
|
protoSocket.AddRequestListenerTyped[*altv1.CompareBacktestRunsRequest, *altv1.CompareBacktestRunsResponse](&c.Communicator, func(req *altv1.CompareBacktestRunsRequest) (*altv1.CompareBacktestRunsResponse, error) {
|
|
if api.compareBacktestRunsResp != nil {
|
|
return api.compareBacktestRunsResp, nil
|
|
}
|
|
return &altv1.CompareBacktestRunsResponse{}, nil
|
|
})
|
|
protoSocket.AddRequestListenerTyped[*altv1.ImportDailyBarsRequest, *altv1.ImportDailyBarsResponse](&c.Communicator, func(req *altv1.ImportDailyBarsRequest) (*altv1.ImportDailyBarsResponse, error) {
|
|
api.recordCall("import_daily_bars")
|
|
api.setImportDailyBarsReq(req)
|
|
if api.importDailyBarsResp != nil {
|
|
return api.importDailyBarsResp, nil
|
|
}
|
|
return &altv1.ImportDailyBarsResponse{}, nil
|
|
})
|
|
protoSocket.AddRequestListenerTyped[*altv1.StartPaperTradingRequest, *altv1.StartPaperTradingResponse](&c.Communicator, func(req *altv1.StartPaperTradingRequest) (*altv1.StartPaperTradingResponse, error) {
|
|
api.recordCall("start_paper_trading")
|
|
api.setStartPaperReq(req)
|
|
if api.startPaperResp != nil {
|
|
return api.startPaperResp, nil
|
|
}
|
|
return &altv1.StartPaperTradingResponse{}, nil
|
|
})
|
|
protoSocket.AddRequestListenerTyped[*altv1.GetPaperTradingStateRequest, *altv1.GetPaperTradingStateResponse](&c.Communicator, func(req *altv1.GetPaperTradingStateRequest) (*altv1.GetPaperTradingStateResponse, error) {
|
|
api.recordCall("get_paper_trading_state")
|
|
api.setGetPaperStateReq(req)
|
|
if api.getPaperStateResp != nil {
|
|
return api.getPaperStateResp, nil
|
|
}
|
|
return &altv1.GetPaperTradingStateResponse{}, nil
|
|
})
|
|
}
|
|
|
|
if err := srv.Start(ctx); err != nil {
|
|
t.Fatalf("start fake API server: %v", err)
|
|
}
|
|
t.Cleanup(func() { _ = srv.Stop() })
|
|
|
|
return fmt.Sprintf("ws://%s:%d%s", host, port, path)
|
|
}
|
|
|
|
func freePort(t *testing.T) int {
|
|
t.Helper()
|
|
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
|
if err != nil {
|
|
t.Fatalf("reserve TCP port: %v", err)
|
|
}
|
|
defer listener.Close()
|
|
addr, ok := listener.Addr().(*net.TCPAddr)
|
|
if !ok {
|
|
t.Fatalf("unexpected listener address type %T", listener.Addr())
|
|
}
|
|
return addr.Port
|
|
}
|
|
|
|
func dialTestClient(t *testing.T, url string) *APIClient {
|
|
t.Helper()
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
t.Cleanup(cancel)
|
|
client, err := Dial(ctx, url)
|
|
if err != nil {
|
|
t.Fatalf("dial fake API: %v", err)
|
|
}
|
|
t.Cleanup(func() { _ = client.Close() })
|
|
return client
|
|
}
|
|
|
|
func TestAPIClientHello(t *testing.T) {
|
|
url := startFakeAPIServer(t, &fakeAPI{helloCaps: []string{"hello", "market-read"}})
|
|
client := dialTestClient(t, url)
|
|
|
|
resp, err := client.Hello(context.Background(), &altv1.HelloRequest{AltProtocolVersion: "alt.v1"})
|
|
if err != nil {
|
|
t.Fatalf("hello: %v", err)
|
|
}
|
|
if resp.GetServerName() != "alt-api-fake" {
|
|
t.Errorf("server name = %q, want alt-api-fake", resp.GetServerName())
|
|
}
|
|
if resp.GetAltProtocolVersion() != "alt.v1" {
|
|
t.Errorf("protocol version = %q, want alt.v1", resp.GetAltProtocolVersion())
|
|
}
|
|
}
|
|
|
|
func TestAPIClientListInstruments(t *testing.T) {
|
|
api := &fakeAPI{instResp: &altv1.ListInstrumentsResponse{
|
|
Instruments: []*altv1.Instrument{{Id: "KRX:005930", Symbol: "005930"}},
|
|
}}
|
|
url := startFakeAPIServer(t, api)
|
|
client := dialTestClient(t, url)
|
|
|
|
resp, err := client.ListInstruments(context.Background(), &altv1.ListInstrumentsRequest{Market: altv1.Market_MARKET_KR})
|
|
if err != nil {
|
|
t.Fatalf("list instruments: %v", err)
|
|
}
|
|
if resp.GetError() != nil {
|
|
t.Fatalf("unexpected typed error: %+v", resp.GetError())
|
|
}
|
|
if len(resp.GetInstruments()) != 1 {
|
|
t.Fatalf("instruments = %d, want 1", len(resp.GetInstruments()))
|
|
}
|
|
if resp.GetInstruments()[0].GetId() != "KRX:005930" {
|
|
t.Errorf("instrument id = %q, want KRX:005930", resp.GetInstruments()[0].GetId())
|
|
}
|
|
|
|
got := api.lastInstReq()
|
|
if got == nil {
|
|
t.Fatal("server did not receive a list instruments request")
|
|
}
|
|
if got.GetMarket() != altv1.Market_MARKET_KR {
|
|
t.Errorf("server received market = %v, want MARKET_KR", got.GetMarket())
|
|
}
|
|
}
|
|
|
|
func TestAPIClientListBars(t *testing.T) {
|
|
api := &fakeAPI{barsResp: &altv1.ListBarsResponse{
|
|
Bars: []*altv1.Bar{{InstrumentId: "KRX:005930", Timeframe: altv1.Timeframe_TIMEFRAME_DAILY}},
|
|
}}
|
|
url := startFakeAPIServer(t, api)
|
|
client := dialTestClient(t, url)
|
|
|
|
req := &altv1.ListBarsRequest{
|
|
InstrumentId: "KRX:005930",
|
|
Timeframe: altv1.Timeframe_TIMEFRAME_DAILY,
|
|
FromUnixMs: 1746057600000,
|
|
ToUnixMs: 1747267200000,
|
|
}
|
|
resp, err := client.ListBars(context.Background(), req)
|
|
if err != nil {
|
|
t.Fatalf("list bars: %v", err)
|
|
}
|
|
if resp.GetError() != nil {
|
|
t.Fatalf("unexpected typed error: %+v", resp.GetError())
|
|
}
|
|
if len(resp.GetBars()) != 1 {
|
|
t.Fatalf("bars = %d, want 1", len(resp.GetBars()))
|
|
}
|
|
|
|
got := api.lastBarsReq()
|
|
if got == nil {
|
|
t.Fatal("server did not receive a list bars request")
|
|
}
|
|
if got.GetInstrumentId() != "KRX:005930" {
|
|
t.Errorf("server received instrument_id = %q, want KRX:005930", got.GetInstrumentId())
|
|
}
|
|
if got.GetTimeframe() != altv1.Timeframe_TIMEFRAME_DAILY {
|
|
t.Errorf("server received timeframe = %v, want TIMEFRAME_DAILY", got.GetTimeframe())
|
|
}
|
|
if got.GetFromUnixMs() != 1746057600000 || got.GetToUnixMs() != 1747267200000 {
|
|
t.Errorf("server received window = [%d,%d], want [1746057600000,1747267200000]", got.GetFromUnixMs(), got.GetToUnixMs())
|
|
}
|
|
}
|
|
|
|
func TestAPIClientConnectUnavailable(t *testing.T) {
|
|
// Reserve then release a port so the dial target is almost certainly closed.
|
|
port := freePort(t)
|
|
url := fmt.Sprintf("ws://127.0.0.1:%d/socket", port)
|
|
|
|
_, err := Dial(context.Background(), url)
|
|
if err == nil {
|
|
t.Fatal("expected dial to a closed port to fail")
|
|
}
|
|
if !errors.Is(err, ErrTransport) {
|
|
t.Errorf("error = %v, want it to wrap ErrTransport", err)
|
|
}
|
|
}
|
|
|
|
func TestAPIClientImportDailyBars(t *testing.T) {
|
|
api := &fakeAPI{importDailyBarsResp: &altv1.ImportDailyBarsResponse{
|
|
Provider: "kis",
|
|
InstrumentCount: 1,
|
|
BarCount: 42,
|
|
}}
|
|
url := startFakeAPIServer(t, api)
|
|
client := dialTestClient(t, url)
|
|
|
|
req := &altv1.ImportDailyBarsRequest{
|
|
Provider: "kis",
|
|
SelectorKind: "watchlist",
|
|
Market: altv1.Market_MARKET_KR,
|
|
Venue: altv1.Venue_VENUE_KRX,
|
|
Symbols: []string{"005930"},
|
|
FromYyyymmdd: "20240527",
|
|
ToYyyymmdd: "20240528",
|
|
}
|
|
resp, err := client.ImportDailyBars(context.Background(), req)
|
|
if err != nil {
|
|
t.Fatalf("ImportDailyBars: %v", err)
|
|
}
|
|
if resp.GetError() != nil {
|
|
t.Fatalf("unexpected typed error: %+v", resp.GetError())
|
|
}
|
|
if resp.GetProvider() != "kis" {
|
|
t.Errorf("resp provider = %q, want kis", resp.GetProvider())
|
|
}
|
|
if resp.GetInstrumentCount() != 1 {
|
|
t.Errorf("resp instrument count = %d, want 1", resp.GetInstrumentCount())
|
|
}
|
|
if resp.GetBarCount() != 42 {
|
|
t.Errorf("resp bar count = %d, want 42", resp.GetBarCount())
|
|
}
|
|
|
|
got := api.lastImportDailyBarsReq()
|
|
if got == nil {
|
|
t.Fatal("server did not receive an import daily bars request")
|
|
}
|
|
if got.GetProvider() != "kis" {
|
|
t.Errorf("server received provider = %q, want kis", got.GetProvider())
|
|
}
|
|
if got.GetSelectorKind() != "watchlist" {
|
|
t.Errorf("server received selector_kind = %q, want watchlist", got.GetSelectorKind())
|
|
}
|
|
if got.GetMarket() != altv1.Market_MARKET_KR {
|
|
t.Errorf("server received market = %v, want MARKET_KR", got.GetMarket())
|
|
}
|
|
if got.GetVenue() != altv1.Venue_VENUE_KRX {
|
|
t.Errorf("server received venue = %v, want VENUE_KRX", got.GetVenue())
|
|
}
|
|
if len(got.GetSymbols()) != 1 || got.GetSymbols()[0] != "005930" {
|
|
t.Errorf("server received symbols = %v, want [005930]", got.GetSymbols())
|
|
}
|
|
if got.GetFromYyyymmdd() != "20240527" {
|
|
t.Errorf("server received from_yyyymmdd = %q, want 20240527", got.GetFromYyyymmdd())
|
|
}
|
|
if got.GetToYyyymmdd() != "20240528" {
|
|
t.Errorf("server received to_yyyymmdd = %q, want 20240528", got.GetToYyyymmdd())
|
|
}
|
|
}
|