package transport_test import ( "context" "fmt" "net" "testing" "time" "go.uber.org/zap" "google.golang.org/protobuf/proto" toki "git.toki-labs.com/toki/proto-socket/go" "iop/apps/node/internal/transport" eventpkg "iop/packages/go/events" iop "iop/proto/gen/iop" ) // startMiniEdge starts a minimal mock edge server that accepts one connection // and responds to RegisterRequest. Returns the server, the accepted edge-side // client channel, and the listen address. func startMiniEdge(t *testing.T, ctx context.Context) (server *toki.TcpServer, acceptedCh chan *toki.TcpClient, addr string) { t.Helper() listenAddr := getFreePort(t) host, portStr, _ := net.SplitHostPort(listenAddr) port := 0 fmt.Sscanf(portStr, "%d", &port) ch := make(chan *toki.TcpClient, 4) srv := toki.NewTcpServer(host, port, func(conn net.Conn) *toki.TcpClient { client := toki.NewTcpClient(conn, 30, 10, edgeParserMap()) toki.AddRequestListenerTyped[*iop.RegisterRequest, *iop.RegisterResponse]( &client.Communicator, func(_ *iop.RegisterRequest) (*iop.RegisterResponse, error) { return &iop.RegisterResponse{ Accepted: true, NodeId: "sess-test-node", Config: &iop.NodeConfigPayload{}, }, nil }, ) ch <- client return client }) if err := srv.Start(ctx); err != nil { t.Fatalf("start mini edge: %v", err) } t.Cleanup(func() { srv.Stop() }) return srv, ch, listenAddr } // TestSessionDoneSignalOnRemoteDisconnect verifies that Done() is closed when // edge closes the connection, and IsLocalShutdown() returns false. func TestSessionDoneSignalOnRemoteDisconnect(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() _, acceptedCh, addr := startMiniEdge(t, ctx) result, err := transport.DialEdge(ctx, addr, "tok", zap.NewNop()) if err != nil { t.Fatalf("dial: %v", err) } edgeClient := waitForAcceptedClient(t, acceptedCh) // Done() must not be closed yet. select { case <-result.Session.Done(): t.Fatal("Done() closed before disconnect") default: } // Remote close → Done() must close. _ = edgeClient.Close() select { case <-result.Session.Done(): case <-time.After(2 * time.Second): t.Fatal("timeout: Done() did not close after remote disconnect") } if result.Session.IsLocalShutdown() { t.Fatal("IsLocalShutdown() must be false for remote disconnect") } } // TestSessionDoneSignalOnLocalClose verifies Done() is closed on local Close() // and IsLocalShutdown() returns true. func TestSessionDoneSignalOnLocalClose(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() _, _, addr := startMiniEdge(t, ctx) result, err := transport.DialEdge(ctx, addr, "tok", zap.NewNop()) if err != nil { t.Fatalf("dial: %v", err) } _ = result.Session.Close() select { case <-result.Session.Done(): case <-time.After(2 * time.Second): t.Fatal("timeout: Done() did not close after local Close()") } if !result.Session.IsLocalShutdown() { t.Fatal("IsLocalShutdown() must be true after local Close()") } } type mockHandler struct { runReqCh chan *iop.RunRequest } func (m *mockHandler) OnRunRequest(ctx context.Context, sess *transport.Session, req *iop.RunRequest) error { m.runReqCh <- req return sess.Send(&iop.RunEvent{RunId: req.GetRunId(), Type: "test_event", NodeId: "test-node"}) } func (m *mockHandler) OnCancel(ctx context.Context, sess *transport.Session, req *iop.CancelRequest) error { return nil } func (m *mockHandler) OnCommandRequest(ctx context.Context, sess *transport.Session, req *iop.NodeCommandRequest) (*iop.NodeCommandResponse, error) { return &iop.NodeCommandResponse{Error: "not implemented in mock"}, nil } func (m *mockHandler) OnConfigRefresh(_ context.Context, _ *transport.Session, req *iop.NodeConfigRefreshRequest) (*iop.NodeConfigRefreshResponse, error) { return &iop.NodeConfigRefreshResponse{ RequestId: req.GetRequestId(), Status: iop.NodeConfigRefreshStatus_NODE_CONFIG_REFRESH_STATUS_RESTART_REQUIRED, }, nil } func (m *mockHandler) OnProviderTunnelRequest(ctx context.Context, sess *transport.Session, req *iop.ProviderTunnelRequest) error { return nil } func getFreePort(t *testing.T) string { l, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("listen: %v", err) } addr := l.Addr().String() l.Close() return addr } func edgeParserMap() toki.ParserMap { return toki.ParserMap{ toki.TypeNameOf(&iop.RunEvent{}): func(b []byte) (proto.Message, error) { m := &iop.RunEvent{} return m, proto.Unmarshal(b, m) }, toki.TypeNameOf(&iop.RegisterRequest{}): func(b []byte) (proto.Message, error) { m := &iop.RegisterRequest{} return m, proto.Unmarshal(b, m) }, } } func waitForAcceptedClient(t *testing.T, acceptedCh <-chan *toki.TcpClient) *toki.TcpClient { t.Helper() select { case client := <-acceptedCh: return client case <-time.After(2 * time.Second): t.Fatal("edge server did not accept connection") return nil } } func TestNodeClientIntegration(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() logger := zap.NewNop() listenAddr := getFreePort(t) host, portStr, _ := net.SplitHostPort(listenAddr) port := 0 fmt.Sscanf(portStr, "%d", &port) // 1. Mock Edge 서버 구동 acceptedCh := make(chan *toki.TcpClient, 1) registerReqCh := make(chan *iop.RegisterRequest, 1) server := toki.NewTcpServer(host, port, func(conn net.Conn) *toki.TcpClient { client := toki.NewTcpClient(conn, 30, 10, edgeParserMap()) toki.AddRequestListenerTyped[*iop.RegisterRequest, *iop.RegisterResponse]( &client.Communicator, func(req *iop.RegisterRequest) (*iop.RegisterResponse, error) { registerReqCh <- req return &iop.RegisterResponse{ Accepted: true, NodeId: "test-node", Alias: "test-alias", Config: &iop.NodeConfigPayload{ Runtime: &iop.NodeRuntimeConfig{Concurrency: 1}, }, }, nil }, ) acceptedCh <- client return client }) if err := server.Start(ctx); err != nil { t.Fatalf("failed to start mock edge server: %v", err) } defer server.Stop() // 2. Node 클라이언트 접속 handler := &mockHandler{ runReqCh: make(chan *iop.RunRequest, 1), } result, err := transport.DialEdge(ctx, listenAddr, "test-token", logger) if err != nil { t.Fatalf("failed to dial edge: %v", err) } defer result.Session.Close() edgeEventCh := make(chan *iop.EdgeNodeEvent, 1) result.Session.SetEventHandler(func(event *iop.EdgeNodeEvent) { edgeEventCh <- event }) result.Session.SetHandler(handler) if result.NodeID != "test-node" { t.Fatalf("expected node id %q, got %q", "test-node", result.NodeID) } if result.Alias != "test-alias" { t.Fatalf("expected alias %q, got %q", "test-alias", result.Alias) } select { case req := <-registerReqCh: if req.GetToken() != "test-token" { t.Fatalf("expected token %q, got %q", "test-token", req.GetToken()) } case <-time.After(2 * time.Second): t.Fatal("timeout waiting for register request") } edgeClient := waitForAcceptedClient(t, acceptedCh) runEventCh := make(chan *iop.RunEvent, 1) toki.AddListenerTyped[*iop.RunEvent](&edgeClient.Communicator, func(event *iop.RunEvent) { runEventCh <- event }) // 3. Edge -> Node 로 RunRequest 전송 runReq := &iop.RunRequest{ RunId: "test-run", Adapter: "test-adapter", } if err := edgeClient.Send(runReq); err != nil { t.Fatalf("failed to send run request: %v", err) } // 4. Node에서 RunRequest 수신 확인 select { case receivedReq := <-handler.runReqCh: if receivedReq.GetRunId() != runReq.GetRunId() { t.Fatalf("expected run id %q, got %q", runReq.GetRunId(), receivedReq.GetRunId()) } case <-time.After(2 * time.Second): t.Fatal("timeout waiting for run request on node handler") } select { case event := <-runEventCh: if event.GetRunId() != runReq.GetRunId() { t.Fatalf("expected run event id %q, got %q", runReq.GetRunId(), event.GetRunId()) } if event.GetType() != "test_event" { t.Fatalf("expected run event type %q, got %q", "test_event", event.GetType()) } case <-time.After(2 * time.Second): t.Fatal("timeout waiting for run event from node session") } if err := edgeClient.Close(); err != nil { t.Fatalf("close edge client: %v", err) } select { case event := <-edgeEventCh: if event.GetType() != eventpkg.TypeEdgeDisconnected { t.Fatalf("event type: got %q want %q", event.GetType(), eventpkg.TypeEdgeDisconnected) } if event.GetNodeId() != "test-node" || event.GetAlias() != "test-alias" || event.GetReason() != eventpkg.ReasonTransportClosed { t.Fatalf("unexpected edge disconnected event: %+v", event) } if event.GetMetadata()[eventpkg.MetadataTransportCloseReason] != toki.DisconnectReasonRemoteClosed { t.Fatalf("transport close reason: got %q want %q", event.GetMetadata()[eventpkg.MetadataTransportCloseReason], toki.DisconnectReasonRemoteClosed) } if event.GetMetadata()[eventpkg.MetadataTransportCloseError] == "" { t.Fatalf("expected transport close error metadata, got %+v", event.GetMetadata()) } case <-time.After(2 * time.Second): t.Fatal("timeout waiting for edge disconnected event") } }