alt/services/api/internal/workerclient/client_test.go
toki ec938475e0 chore: socket rail implementation and task archival
- Add backtest and market socket handlers for API and Worker
- Add proto definitions for backtest and market services
- Update client socket integration and tests
- Generate protobuf code for Go and Dart
- Archive completed agent tasks under agent-task/archive/2026/05/
2026-05-30 22:57:18 +09:00

317 lines
9.9 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_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_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)
}
}