package server import ( "context" "encoding/json" "errors" "net" "net/http" "net/http/httptest" "net/url" "strings" "testing" "testing/synctest" "time" "go.miragespace.co/specter/spec/chord" "go.miragespace.co/specter/spec/mocks" "go.miragespace.co/specter/spec/protocol" "go.miragespace.co/specter/spec/transport" "go.miragespace.co/specter/spec/tun" "github.com/go-chi/chi/v5" "github.com/stretchr/testify/mock" "github.com/stretchr/testify/require" ) type handlerVNode struct { chord.VNode prefixList func(context.Context, []byte) ([][]byte, error) } func (n handlerVNode) PrefixList(ctx context.Context, prefix []byte) ([][]byte, error) { return n.prefixList(ctx, prefix) } type handlerTransport struct { transport.Transport dialStream func(context.Context, *protocol.Node, protocol.Stream_Type) (net.Conn, error) } func (tr handlerTransport) DialStream(ctx context.Context, peer *protocol.Node, kind protocol.Stream_Type) (net.Conn, error) { return tr.dialStream(ctx, peer, kind) } type handlerQueryService func(context.Context, *protocol.ListTunnelsRequest) (*protocol.ListTunnelsResponse, error) func (f handlerQueryService) ListTunnels(ctx context.Context, req *protocol.ListTunnelsRequest) (*protocol.ListTunnelsResponse, error) { return f(ctx, req) } func handlerRouter(node chord.VNode, tr transport.Transport) http.Handler { serv := &Server{Config: Config{Chord: node, TunnelTransport: tr}} router := chi.NewRouter() router.Mount("/clients", TunnelServerHandler(serv)) return router } func handlerClientTransport(t *testing.T, service handlerQueryService, onDial func(context.Context)) transport.Transport { t.Helper() client := httptest.NewServer(protocol.NewClientQueryServiceServer(service)) t.Cleanup(client.Close) return handlerTransport{dialStream: func(ctx context.Context, peer *protocol.Node, kind protocol.Stream_Type) (net.Conn, error) { if peer.GetId() != 111111 || peer.GetAddress() != "fake-address" || !peer.GetRendezvous() || kind != protocol.Stream_RPC { t.Errorf("unexpected client RPC destination: %v, stream %v", peer, kind) } if onDial != nil { onDial(ctx) } return (&net.Dialer{}).DialContext(ctx, "tcp", client.Listener.Addr().String()) }} } func handlerTunnelInfo(t *testing.T, w *httptest.ResponseRecorder) tunnelsInfo { t.Helper() require.Equal(t, http.StatusOK, w.Code) require.Equal(t, "application/json; charset=utf-8", w.Header().Get("Content-Type")) require.Equal(t, "no-store", w.Header().Get("Cache-Control")) var info tunnelsInfo require.NoError(t, json.Unmarshal(w.Body.Bytes(), &info)) require.NotNil(t, info.Tunnels, "empty tunnel results must be arrays") return info } func handlerTunnelRows(t *testing.T, info tunnelsInfo) map[string]clientTunnel { t.Helper() rows := make(map[string]clientTunnel, len(info.Tunnels)) for i, tunnel := range info.Tunnels { if i > 0 { require.Less(t, info.Tunnels[i-1].Hostname, tunnel.Hostname, "tunnels must be sorted and unique") } rows[tunnel.Hostname] = tunnel } return rows } func TestHandlerListConnectedClients(t *testing.T) { node := new(mocks.VNode) clientT := new(mocks.Transport) clientT.On("ListConnected").Return([]transport.ConnectedPeer{ { Identity: &protocol.Node{Id: 222222, Address: "second-client"}, Addr: &net.UDPAddr{IP: net.ParseIP("192.0.2.2"), Port: 4200}, Version: "v2.0.0", }, { Identity: &protocol.Node{Id: 111111, Address: "fake-address"}, Addr: &net.UDPAddr{IP: net.ParseIP("192.0.2.1"), Port: 4200}, Version: "v1.0.0", }, }).Once() clientT.On("Identity").Return(&protocol.Node{Address: "local-node"}).Once() w := httptest.NewRecorder() handlerRouter(node, clientT).ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/clients", nil)) require.Equal(t, http.StatusOK, w.Code) require.Equal(t, "application/json; charset=utf-8", w.Header().Get("Content-Type")) require.Equal(t, "no-store", w.Header().Get("Cache-Control")) var body map[string]json.RawMessage require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body)) var observedAt string require.NoError(t, json.Unmarshal(body["observedAt"], &observedAt)) _, err := time.Parse(time.RFC3339, observedAt) require.NoError(t, err) delete(body, "observedAt") payload, err := json.Marshal(body) require.NoError(t, err) require.JSONEq(t, `{ "node": "local-node", "clients": [ {"clientId":"111111","identity":"111111/fake-address","address":"192.0.2.1:4200","version":"v1.0.0","url":"/_internal/tun/111111/fake-address"}, {"clientId":"222222","identity":"222222/second-client","address":"192.0.2.2:4200","version":"v2.0.0","url":"/_internal/tun/222222/second-client"} ] }`, string(payload)) require.Empty(t, node.Calls, "the connected-client list must use only the local transport snapshot") node.AssertExpectations(t) clientT.AssertExpectations(t) } func TestHandlerEmptyConnectedClientsIsArray(t *testing.T) { clientT := new(mocks.Transport) clientT.On("ListConnected").Return([]transport.ConnectedPeer(nil)).Once() clientT.On("Identity").Return(&protocol.Node{Address: "local-node"}).Once() w := httptest.NewRecorder() handlerRouter(new(mocks.VNode), clientT).ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/clients", nil)) var body map[string]json.RawMessage require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body)) require.JSONEq(t, `[]`, string(body["clients"])) clientT.AssertExpectations(t) } func TestHandlerClientIdentityURLRoundTrip(t *testing.T) { for _, address := range []string{`client/?x=1&y=2`, "client%2Fescaped"} { t.Run(address, func(t *testing.T) { clientT := new(mocks.Transport) clientT.On("ListConnected").Return([]transport.ConnectedPeer{{ Identity: &protocol.Node{Id: 111111, Address: address}, Addr: &net.UDPAddr{IP: net.ParseIP("192.0.2.1"), Port: 4200}, Version: ``, }}).Once() clientT.On("Identity").Return(&protocol.Node{Address: ``}).Once() w := httptest.NewRecorder() handlerRouter(new(mocks.VNode), clientT).ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/clients", nil)) var list connectedInfo require.NoError(t, json.Unmarshal(w.Body.Bytes(), &list)) require.Len(t, list.Clients, 1) require.Equal(t, ``, list.Node) require.Equal(t, ``, list.Clients[0].Version) require.Equal(t, "111111/"+address, list.Clients[0].Identity) require.Equal(t, "/_internal/tun/111111/"+url.PathEscape(address), list.Clients[0].URL) require.NotContains(t, w.Body.String(), "