Skip to content
File

Blob: tun/server/connection_test.go

go161 lines
1package server
2 
3import (
4 "context"
5 "io"
6 "net"
7 "testing"
8 "time"
9 
10 "go.miragespace.co/specter/spec/protocol"
11 "go.miragespace.co/specter/spec/rpc"
12 "go.miragespace.co/specter/spec/tun"
13 
14 "github.com/stretchr/testify/mock"
15 "github.com/stretchr/testify/require"
16)
17 
18type negotiationConn struct {
19 net.Conn
20 writeStarted chan struct{}
21 deadline time.Time
22}
23 
24func (c *negotiationConn) Write(p []byte) (int, error) {
25 select {
26 case c.writeStarted <- struct{}{}:
27 default:
28 }
29 return c.Conn.Write(p)
30}
31 
32func (c *negotiationConn) SetDeadline(deadline time.Time) error {
33 c.deadline = deadline
34 return c.Conn.SetDeadline(deadline)
35}
36 
37func connectionFixture(t *testing.T, direct bool, conn net.Conn) (*Server, *protocol.TunnelRoute, *protocol.Link, <-chan context.Context) {
38 t.Helper()
39 _, _, clientT, chordT, serv := getFixture(t, require.New(t))
40 t.Cleanup(serv.routeCache.Close)
41 t.Cleanup(serv.keylessCache.Close)
42 client, chordNode, tunnelNode := getIdentities()
43 link := &protocol.Link{Hostname: "negotiation.example.com", Alpn: protocol.Link_HTTP}
44 route := &protocol.TunnelRoute{
45 Hostname: link.GetHostname(), ClientDestination: client,
46 ChordDestination: chordNode, TunnelDestination: tunnelNode,
47 }
48 dialContexts := make(chan context.Context, 1)
49 recordDialContext := func(args mock.Arguments) { dialContexts <- args.Get(0).(context.Context) }
50 if direct {
51 clientT.On("Identity").Return(tunnelNode)
52 clientT.On("DialStream", mock.Anything, client, protocol.Stream_DIRECT).Return(conn, nil).Once().Run(recordDialContext)
53 } else {
54 clientT.On("Identity").Return(&protocol.Node{Address: "local-tunnel:123"})
55 chordT.On("DialStream", mock.Anything, chordNode, protocol.Stream_PROXY).Return(conn, nil).Once().Run(recordDialContext)
56 }
57 t.Cleanup(func() {
58 clientT.AssertExpectations(t)
59 chordT.AssertExpectations(t)
60 })
61 return serv, route, link, dialContexts
62}
63 
64func sendProxyStatus(peer net.Conn, status *protocol.TunnelStatus) error {
65 if err := rpc.BoundedReceive(peer, &protocol.TunnelRoute{}, 2048); err != nil {
66 return err
67 }
68 return rpc.Send(peer, status)
69}
70 
71func TestGetConnCancellationInterruptsNegotiationWrites(t *testing.T) {
72 for _, tc := range []struct {
73 name string
74 direct bool
75 }{
76 {name: "proxy_route"},
77 {name: "direct_link", direct: true},
78 } {
79 t.Run(tc.name, func(t *testing.T) {
80 local, peer := net.Pipe()
81 conn := &negotiationConn{
82 Conn: local, writeStarted: make(chan struct{}, 1),
83 }
84 t.Cleanup(func() { conn.Close(); peer.Close() })
85 serv, route, link, _ := connectionFixture(t, tc.direct, conn)
86 ctx, cancel := context.WithCancel(t.Context())
87 defer cancel()
88 done := make(chan error, 1)
89 go func() {
90 _, err := serv.getConn(ctx, route, link)
91 done <- err
92 }()
93 awaitRouteResult(t, conn.writeStarted)
94 require.NoError(t, peer.SetReadDeadline(time.Now().Add(time.Second)))
95 cancel()
96 require.Error(t, awaitRouteResult(t, done))
97 _, err := peer.Read(make([]byte, 1))
98 require.ErrorIs(t, err, io.EOF, "cancelled negotiation left the stream open")
99 })
100 }
101}
102 
103func TestGetConnClosesRejectedProxyStreams(t *testing.T) {
104 local, peer := net.Pipe()
105 t.Cleanup(func() { local.Close(); peer.Close() })
106 serv, route, link, _ := connectionFixture(t, false, local)
107 statusSent := make(chan error, 1)
108 go func() {
109 statusSent <- sendProxyStatus(peer, &protocol.TunnelStatus{Status: protocol.TunnelStatusCode_NO_DIRECT})
110 }()
111 require.NoError(t, peer.SetReadDeadline(time.Now().Add(time.Second)))
112 got, err := serv.getConn(t.Context(), route, link)
113 require.Nil(t, got)
114 require.ErrorIs(t, err, tun.ErrTunnelClientNotConnected)
115 require.NoError(t, awaitRouteResult(t, statusSent))
116 _, err = peer.Read(make([]byte, 1))
117 require.ErrorIs(t, err, io.EOF, "rejected proxy negotiation left the stream open")
118}
119 
120func TestGetConnClearsNegotiationDeadline(t *testing.T) {
121 for _, direct := range []bool{true, false} {
122 name := "proxy"
123 if direct {
124 name = "direct"
125 }
126 t.Run(name, func(t *testing.T) {
127 local, peer := net.Pipe()
128 conn := &negotiationConn{Conn: local}
129 t.Cleanup(func() { conn.Close(); peer.Close() })
130 serv, route, link, dialContexts := connectionFixture(t, direct, conn)
131 peerDone := make(chan error, 1)
132 go func() {
133 if !direct {
134 if err := sendProxyStatus(peer, &protocol.TunnelStatus{}); err != nil {
135 peerDone <- err
136 return
137 }
138 }
139 if err := rpc.BoundedReceive(peer, &protocol.Link{}, 1024); err != nil {
140 peerDone <- err
141 return
142 }
143 _, err := io.ReadFull(peer, make([]byte, 1))
144 peerDone <- err
145 }()
146 ctx, cancel := context.WithCancel(t.Context())
147 defer cancel()
148 got, err := serv.getConn(ctx, route, link)
149 require.NoError(t, err)
150 require.Same(t, conn, got)
151 require.NoError(t, (<-dialContexts).Err(), "completed negotiation cancelled the transport connection context")
152 require.True(t, conn.deadline.IsZero(), "successful stream retained its negotiation deadline")
153 cancel()
154 require.NoError(t, got.SetWriteDeadline(time.Now().Add(time.Second)))
155 _, err = got.Write([]byte{1})
156 require.NoError(t, err, "completed negotiation retained its cancellation hook")
157 require.NoError(t, <-peerDone)
158 })
159 }
160}