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(), "