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