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_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) } }