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 aggregateMonthlyResp *altv1.AggregateMonthlyBarsResponse startPaperResp *altv1.StartPaperTradingResponse getPaperStateResp *altv1.GetPaperTradingStateResponse submitPaperOrderResp *altv1.SubmitPaperOrderResponse cancelPaperOrderResp *altv1.CancelPaperOrderResponse fillPaperOrderResp *altv1.FillPaperOrderResponse submitLiveOrderResp *altv1.SubmitLiveOrderResponse cancelLiveOrderResp *altv1.CancelLiveOrderResponse getLiveOrderResp *altv1.GetLiveOrderResponse getLiveRiskPolicyResp *altv1.GetLiveRiskPolicyResponse getLiveKillSwitchResp *altv1.GetLiveKillSwitchResponse setLiveKillSwitchResp *altv1.SetLiveKillSwitchResponse syncLiveAccountResp *altv1.SyncLiveAccountResponse getLiveAccountSnapResp *altv1.GetLiveAccountSnapshotResponse listLiveAuditEventsResp *altv1.ListLiveAuditEventsResponse schedulerRefreshStatusResp *altv1.SchedulerRefreshStatusResponse schedulerRefreshStatusFunc func(scheduleName string) *altv1.SchedulerRefreshStatusResponse startBacktestFunc func(req *altv1.StartBacktestRequest) (*altv1.StartBacktestResponse, error) getBacktestRunFunc func(req *altv1.GetBacktestRunRequest) (*altv1.GetBacktestRunResponse, error) barsRespFunc func(req *altv1.ListBarsRequest) (*altv1.ListBarsResponse, error) mu sync.Mutex instReq *altv1.ListInstrumentsRequest barsReq *altv1.ListBarsRequest startBacktestReq *altv1.StartBacktestRequest getBacktestRunReq *altv1.GetBacktestRunRequest importDailyBarsReq *altv1.ImportDailyBarsRequest aggregateMonthlyReq *altv1.AggregateMonthlyBarsRequest startPaperReq *altv1.StartPaperTradingRequest getPaperStateReq *altv1.GetPaperTradingStateRequest submitPaperOrderReq *altv1.SubmitPaperOrderRequest cancelPaperOrderReq *altv1.CancelPaperOrderRequest fillPaperOrderReq *altv1.FillPaperOrderRequest submitLiveOrderReq *altv1.SubmitLiveOrderRequest cancelLiveOrderReq *altv1.CancelLiveOrderRequest getLiveOrderReq *altv1.GetLiveOrderRequest getLiveRiskPolicyReq *altv1.GetLiveRiskPolicyRequest getLiveKillSwitchReq *altv1.GetLiveKillSwitchRequest setLiveKillSwitchReq *altv1.SetLiveKillSwitchRequest syncLiveAccountReq *altv1.SyncLiveAccountRequest getLiveAccountSnapReq *altv1.GetLiveAccountSnapshotRequest listLiveAuditEventsReq *altv1.ListLiveAuditEventsRequest schedulerRefreshStatusReq *altv1.SchedulerRefreshStatusRequest 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) setAggregateMonthlyReq(req *altv1.AggregateMonthlyBarsRequest) { f.mu.Lock() defer f.mu.Unlock() f.aggregateMonthlyReq = req } func (f *fakeAPI) lastAggregateMonthlyReq() *altv1.AggregateMonthlyBarsRequest { f.mu.Lock() defer f.mu.Unlock() return f.aggregateMonthlyReq } 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 } func (f *fakeAPI) setSubmitPaperOrderReq(req *altv1.SubmitPaperOrderRequest) { f.mu.Lock() defer f.mu.Unlock() f.submitPaperOrderReq = req } func (f *fakeAPI) lastSubmitPaperOrderReq() *altv1.SubmitPaperOrderRequest { f.mu.Lock() defer f.mu.Unlock() return f.submitPaperOrderReq } func (f *fakeAPI) setCancelPaperOrderReq(req *altv1.CancelPaperOrderRequest) { f.mu.Lock() defer f.mu.Unlock() f.cancelPaperOrderReq = req } func (f *fakeAPI) lastCancelPaperOrderReq() *altv1.CancelPaperOrderRequest { f.mu.Lock() defer f.mu.Unlock() return f.cancelPaperOrderReq } func (f *fakeAPI) setFillPaperOrderReq(req *altv1.FillPaperOrderRequest) { f.mu.Lock() defer f.mu.Unlock() f.fillPaperOrderReq = req } func (f *fakeAPI) lastFillPaperOrderReq() *altv1.FillPaperOrderRequest { f.mu.Lock() defer f.mu.Unlock() return f.fillPaperOrderReq } func (f *fakeAPI) setSubmitLiveOrderReq(req *altv1.SubmitLiveOrderRequest) { f.mu.Lock() defer f.mu.Unlock() f.submitLiveOrderReq = req } func (f *fakeAPI) lastSubmitLiveOrderReq() *altv1.SubmitLiveOrderRequest { f.mu.Lock() defer f.mu.Unlock() return f.submitLiveOrderReq } func (f *fakeAPI) setCancelLiveOrderReq(req *altv1.CancelLiveOrderRequest) { f.mu.Lock() defer f.mu.Unlock() f.cancelLiveOrderReq = req } func (f *fakeAPI) lastCancelLiveOrderReq() *altv1.CancelLiveOrderRequest { f.mu.Lock() defer f.mu.Unlock() return f.cancelLiveOrderReq } func (f *fakeAPI) setGetLiveOrderReq(req *altv1.GetLiveOrderRequest) { f.mu.Lock() defer f.mu.Unlock() f.getLiveOrderReq = req } func (f *fakeAPI) lastGetLiveOrderReq() *altv1.GetLiveOrderRequest { f.mu.Lock() defer f.mu.Unlock() return f.getLiveOrderReq } func (f *fakeAPI) setGetLiveRiskPolicyReq(req *altv1.GetLiveRiskPolicyRequest) { f.mu.Lock() defer f.mu.Unlock() f.getLiveRiskPolicyReq = req } func (f *fakeAPI) lastGetLiveRiskPolicyReq() *altv1.GetLiveRiskPolicyRequest { f.mu.Lock() defer f.mu.Unlock() return f.getLiveRiskPolicyReq } func (f *fakeAPI) setGetLiveKillSwitchReq(req *altv1.GetLiveKillSwitchRequest) { f.mu.Lock() defer f.mu.Unlock() f.getLiveKillSwitchReq = req } func (f *fakeAPI) lastGetLiveKillSwitchReq() *altv1.GetLiveKillSwitchRequest { f.mu.Lock() defer f.mu.Unlock() return f.getLiveKillSwitchReq } func (f *fakeAPI) setSetLiveKillSwitchReq(req *altv1.SetLiveKillSwitchRequest) { f.mu.Lock() defer f.mu.Unlock() f.setLiveKillSwitchReq = req } func (f *fakeAPI) lastSetLiveKillSwitchReq() *altv1.SetLiveKillSwitchRequest { f.mu.Lock() defer f.mu.Unlock() return f.setLiveKillSwitchReq } func (f *fakeAPI) setSyncLiveAccountReq(req *altv1.SyncLiveAccountRequest) { f.mu.Lock() defer f.mu.Unlock() f.syncLiveAccountReq = req } func (f *fakeAPI) lastSyncLiveAccountReq() *altv1.SyncLiveAccountRequest { f.mu.Lock() defer f.mu.Unlock() return f.syncLiveAccountReq } func (f *fakeAPI) setGetLiveAccountSnapReq(req *altv1.GetLiveAccountSnapshotRequest) { f.mu.Lock() defer f.mu.Unlock() f.getLiveAccountSnapReq = req } func (f *fakeAPI) lastGetLiveAccountSnapReq() *altv1.GetLiveAccountSnapshotRequest { f.mu.Lock() defer f.mu.Unlock() return f.getLiveAccountSnapReq } func (f *fakeAPI) setListLiveAuditEventsReq(req *altv1.ListLiveAuditEventsRequest) { f.mu.Lock() defer f.mu.Unlock() f.listLiveAuditEventsReq = req } func (f *fakeAPI) lastListLiveAuditEventsReq() *altv1.ListLiveAuditEventsRequest { f.mu.Lock() defer f.mu.Unlock() return f.listLiveAuditEventsReq } func (f *fakeAPI) setSchedulerRefreshStatusReq(req *altv1.SchedulerRefreshStatusRequest) { f.mu.Lock() defer f.mu.Unlock() f.schedulerRefreshStatusReq = req } func (f *fakeAPI) lastSchedulerRefreshStatusReq() *altv1.SchedulerRefreshStatusRequest { f.mu.Lock() defer f.mu.Unlock() return f.schedulerRefreshStatusReq } // 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.barsRespFunc != nil { return api.barsRespFunc(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.recordCall("start_backtest") api.setStartBacktestReq(req) if api.startBacktestFunc != nil { return api.startBacktestFunc(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.AggregateMonthlyBarsRequest, *altv1.AggregateMonthlyBarsResponse](&c.Communicator, func(req *altv1.AggregateMonthlyBarsRequest) (*altv1.AggregateMonthlyBarsResponse, error) { api.recordCall("aggregate_monthly_bars") api.setAggregateMonthlyReq(req) if api.aggregateMonthlyResp != nil { return api.aggregateMonthlyResp, nil } return &altv1.AggregateMonthlyBarsResponse{}, 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 }) protoSocket.AddRequestListenerTyped[*altv1.SubmitPaperOrderRequest, *altv1.SubmitPaperOrderResponse](&c.Communicator, func(req *altv1.SubmitPaperOrderRequest) (*altv1.SubmitPaperOrderResponse, error) { api.recordCall("submit_paper_order") api.setSubmitPaperOrderReq(req) if api.submitPaperOrderResp != nil { return api.submitPaperOrderResp, nil } return &altv1.SubmitPaperOrderResponse{}, nil }) protoSocket.AddRequestListenerTyped[*altv1.CancelPaperOrderRequest, *altv1.CancelPaperOrderResponse](&c.Communicator, func(req *altv1.CancelPaperOrderRequest) (*altv1.CancelPaperOrderResponse, error) { api.recordCall("cancel_paper_order") api.setCancelPaperOrderReq(req) if api.cancelPaperOrderResp != nil { return api.cancelPaperOrderResp, nil } return &altv1.CancelPaperOrderResponse{}, nil }) protoSocket.AddRequestListenerTyped[*altv1.FillPaperOrderRequest, *altv1.FillPaperOrderResponse](&c.Communicator, func(req *altv1.FillPaperOrderRequest) (*altv1.FillPaperOrderResponse, error) { api.recordCall("fill_paper_order") api.setFillPaperOrderReq(req) if api.fillPaperOrderResp != nil { return api.fillPaperOrderResp, nil } return &altv1.FillPaperOrderResponse{}, nil }) protoSocket.AddRequestListenerTyped[*altv1.SubmitLiveOrderRequest, *altv1.SubmitLiveOrderResponse](&c.Communicator, func(req *altv1.SubmitLiveOrderRequest) (*altv1.SubmitLiveOrderResponse, error) { api.recordCall("submit_live_order") api.setSubmitLiveOrderReq(req) if api.submitLiveOrderResp != nil { return api.submitLiveOrderResp, nil } return &altv1.SubmitLiveOrderResponse{}, nil }) protoSocket.AddRequestListenerTyped[*altv1.CancelLiveOrderRequest, *altv1.CancelLiveOrderResponse](&c.Communicator, func(req *altv1.CancelLiveOrderRequest) (*altv1.CancelLiveOrderResponse, error) { api.recordCall("cancel_live_order") api.setCancelLiveOrderReq(req) if api.cancelLiveOrderResp != nil { return api.cancelLiveOrderResp, nil } return &altv1.CancelLiveOrderResponse{}, nil }) protoSocket.AddRequestListenerTyped[*altv1.GetLiveOrderRequest, *altv1.GetLiveOrderResponse](&c.Communicator, func(req *altv1.GetLiveOrderRequest) (*altv1.GetLiveOrderResponse, error) { api.recordCall("get_live_order") api.setGetLiveOrderReq(req) if api.getLiveOrderResp != nil { return api.getLiveOrderResp, nil } return &altv1.GetLiveOrderResponse{}, nil }) protoSocket.AddRequestListenerTyped[*altv1.GetLiveRiskPolicyRequest, *altv1.GetLiveRiskPolicyResponse](&c.Communicator, func(req *altv1.GetLiveRiskPolicyRequest) (*altv1.GetLiveRiskPolicyResponse, error) { api.recordCall("get_live_risk_policy") api.setGetLiveRiskPolicyReq(req) if api.getLiveRiskPolicyResp != nil { return api.getLiveRiskPolicyResp, nil } return &altv1.GetLiveRiskPolicyResponse{}, nil }) protoSocket.AddRequestListenerTyped[*altv1.GetLiveKillSwitchRequest, *altv1.GetLiveKillSwitchResponse](&c.Communicator, func(req *altv1.GetLiveKillSwitchRequest) (*altv1.GetLiveKillSwitchResponse, error) { api.recordCall("get_live_kill_switch") api.setGetLiveKillSwitchReq(req) if api.getLiveKillSwitchResp != nil { return api.getLiveKillSwitchResp, nil } return &altv1.GetLiveKillSwitchResponse{}, nil }) protoSocket.AddRequestListenerTyped[*altv1.SetLiveKillSwitchRequest, *altv1.SetLiveKillSwitchResponse](&c.Communicator, func(req *altv1.SetLiveKillSwitchRequest) (*altv1.SetLiveKillSwitchResponse, error) { api.recordCall("set_live_kill_switch") api.setSetLiveKillSwitchReq(req) if api.setLiveKillSwitchResp != nil { return api.setLiveKillSwitchResp, nil } return &altv1.SetLiveKillSwitchResponse{}, nil }) protoSocket.AddRequestListenerTyped[*altv1.SyncLiveAccountRequest, *altv1.SyncLiveAccountResponse](&c.Communicator, func(req *altv1.SyncLiveAccountRequest) (*altv1.SyncLiveAccountResponse, error) { api.recordCall("sync_live_account") api.setSyncLiveAccountReq(req) if api.syncLiveAccountResp != nil { return api.syncLiveAccountResp, nil } return &altv1.SyncLiveAccountResponse{}, nil }) protoSocket.AddRequestListenerTyped[*altv1.GetLiveAccountSnapshotRequest, *altv1.GetLiveAccountSnapshotResponse](&c.Communicator, func(req *altv1.GetLiveAccountSnapshotRequest) (*altv1.GetLiveAccountSnapshotResponse, error) { api.recordCall("get_live_account_snapshot") api.setGetLiveAccountSnapReq(req) if api.getLiveAccountSnapResp != nil { return api.getLiveAccountSnapResp, nil } return &altv1.GetLiveAccountSnapshotResponse{}, nil }) protoSocket.AddRequestListenerTyped[*altv1.ListLiveAuditEventsRequest, *altv1.ListLiveAuditEventsResponse](&c.Communicator, func(req *altv1.ListLiveAuditEventsRequest) (*altv1.ListLiveAuditEventsResponse, error) { api.recordCall("list_live_audit_events") api.setListLiveAuditEventsReq(req) if api.listLiveAuditEventsResp != nil { return api.listLiveAuditEventsResp, nil } return &altv1.ListLiveAuditEventsResponse{}, nil }) protoSocket.AddRequestListenerTyped[*altv1.SchedulerRefreshStatusRequest, *altv1.SchedulerRefreshStatusResponse](&c.Communicator, func(req *altv1.SchedulerRefreshStatusRequest) (*altv1.SchedulerRefreshStatusResponse, error) { api.recordCall("scheduler_refresh_status") api.setSchedulerRefreshStatusReq(req) if api.schedulerRefreshStatusFunc != nil { if resp := api.schedulerRefreshStatusFunc(req.GetScheduleName()); resp != nil { return resp, nil } } if api.schedulerRefreshStatusResp != nil { return api.schedulerRefreshStatusResp, nil } return &altv1.SchedulerRefreshStatusResponse{}, 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()) } } func TestAPIClientStartBacktest(t *testing.T) { api := &fakeAPI{startBacktestResp: &altv1.StartBacktestResponse{ Run: &altv1.BacktestRun{Id: "run-1"}, }} url := startFakeAPIServer(t, api) client := dialTestClient(t, url) req := &altv1.StartBacktestRequest{ Spec: &altv1.BacktestRunSpec{ StrategyId: "strat-abc", Selector: &altv1.BacktestInputSelector{ InstrumentIds: []string{"005930"}, Symbols: []string{"AAPL"}, }, }, } resp, err := client.StartBacktest(context.Background(), req) if err != nil { t.Fatalf("StartBacktest: %v", err) } if resp.GetError() != nil { t.Fatalf("unexpected typed error: %+v", resp.GetError()) } if resp.GetRun().GetId() != "run-1" { t.Errorf("resp run id = %q, want run-1", resp.GetRun().GetId()) } got := api.lastStartBacktestReq() if got == nil { t.Fatal("server did not receive a start backtest request") } if got.GetSpec().GetStrategyId() != "strat-abc" { t.Errorf("server received strategy_id = %q, want strat-abc", got.GetSpec().GetStrategyId()) } pSel := got.GetSpec().GetSelector() if pSel == nil { t.Fatal("server received nil selector") } if len(pSel.GetInstrumentIds()) != 1 || pSel.GetInstrumentIds()[0] != "005930" { t.Errorf("server received instrument_ids = %v, want [005930]", pSel.GetInstrumentIds()) } if len(pSel.GetSymbols()) != 1 || pSel.GetSymbols()[0] != "AAPL" { t.Errorf("server received symbols = %v, want [AAPL]", pSel.GetSymbols()) } }