File
Blob: tun/client/acme_test.go
| 1 | package client |
| 2 | |
| 3 | import ( |
| 4 | "bytes" |
| 5 | "crypto/ed25519" |
| 6 | "fmt" |
| 7 | "io" |
| 8 | "net" |
| 9 | "net/http" |
| 10 | "os" |
| 11 | "testing" |
| 12 | |
| 13 | "go.miragespace.co/specter/spec/acme" |
| 14 | "go.miragespace.co/specter/spec/chord" |
| 15 | "go.miragespace.co/specter/spec/mocks" |
| 16 | "go.miragespace.co/specter/spec/pki" |
| 17 | "go.miragespace.co/specter/spec/pow" |
| 18 | "go.miragespace.co/specter/spec/protocol" |
| 19 | |
| 20 | "github.com/stretchr/testify/mock" |
| 21 | "github.com/stretchr/testify/require" |
| 22 | "github.com/zhangyunhao116/skipmap" |
| 23 | "go.uber.org/zap/zaptest" |
| 24 | ) |
| 25 | |
| 26 | func TestAcmeInstruction(t *testing.T) { |
| 27 | as := require.New(t) |
| 28 | logger := zaptest.NewLogger(t) |
| 29 | |
| 30 | file, err := os.CreateTemp("", "client") |
| 31 | as.NoError(err) |
| 32 | defer os.Remove(file.Name()) |
| 33 | |
| 34 | ctx := t.Context() |
| 35 | |
| 36 | token := &protocol.ClientToken{ |
| 37 | Token: []byte("test"), |
| 38 | } |
| 39 | cl := &protocol.Node{ |
| 40 | Id: chord.Random(), |
| 41 | } |
| 42 | |
| 43 | hostname := "custom.domain.com" |
| 44 | acmeName := "acme.example.com" |
| 45 | acmeContent := "12345678" |
| 46 | |
| 47 | der, cert, key := makeCertificate(as, logger, cl, token, nil) |
| 48 | privKey, err := pki.UnmarshalPrivateKey([]byte(key)) |
| 49 | as.NoError(err) |
| 50 | |
| 51 | cfg := &Config{ |
| 52 | path: file.Name(), |
| 53 | router: skipmap.NewString[route](), |
| 54 | Apex: testApex, |
| 55 | Certificate: cert, |
| 56 | PrivKey: key, |
| 57 | Tunnels: []Tunnel{ |
| 58 | { |
| 59 | Target: "tcp://127.0.0.1:2345", |
| 60 | }, |
| 61 | }, |
| 62 | } |
| 63 | as.NoError(cfg.validate()) |
| 64 | |
| 65 | m := func(s *mocks.TunnelService, t1 *mocks.MemoryTransport, publishCall *mock.Call) { |
| 66 | s.On("AcmeInstruction", mock.Anything, mock.MatchedBy(func(req *protocol.InstructionRequest) bool { |
| 67 | return powValidateFunc(hostname, privKey)(req.GetProof(), req.GetHostname()) |
| 68 | })).Return(&protocol.InstructionResponse{ |
| 69 | Name: acmeName, |
| 70 | Content: acmeContent, |
| 71 | }, nil) |
| 72 | |
| 73 | defaultNoHostnames(s) |
| 74 | transportHelper(t1, der) |
| 75 | } |
| 76 | |
| 77 | listenCfg := &net.ListenConfig{} |
| 78 | sListener, err := listenCfg.Listen(ctx, "tcp", "127.0.0.1:0") |
| 79 | as.NoError(err) |
| 80 | defer sListener.Close() |
| 81 | |
| 82 | client, _, assertion := setupClient(t, as, ctx, logger, nil, cfg, nil, m, false, 1) |
| 83 | defer assertion() |
| 84 | defer client.Close() |
| 85 | |
| 86 | client.ServerListener = sListener |
| 87 | |
| 88 | client.Start(ctx) |
| 89 | |
| 90 | resp, err := client.GetAcmeInstruction(ctx, hostname) |
| 91 | as.NoError(err) |
| 92 | as.NotNil(resp) |
| 93 | as.EqualValues(acmeName, resp.GetName()) |
| 94 | as.EqualValues(acmeContent, resp.GetContent()) |
| 95 | |
| 96 | c := &http.Client{ |
| 97 | Timeout: acme.HashcashExpires, |
| 98 | } |
| 99 | httpReq, err := http.NewRequest(http.MethodGet, fmt.Sprintf("http://%s/api/acme/%s", sListener.Addr().String(), hostname), nil) |
| 100 | as.NoError(err) |
| 101 | httpResp, err := c.Do(httpReq) |
| 102 | as.NoError(err) |
| 103 | defer httpResp.Body.Close() |
| 104 | |
| 105 | body, err := io.ReadAll(httpResp.Body) |
| 106 | as.NoError(err) |
| 107 | as.Contains(string(body), resp.GetContent()) |
| 108 | } |
| 109 | |
| 110 | func TestAcmeValidation(t *testing.T) { |
| 111 | as := require.New(t) |
| 112 | logger := zaptest.NewLogger(t) |
| 113 | |
| 114 | file, err := os.CreateTemp("", "client") |
| 115 | as.NoError(err) |
| 116 | defer os.Remove(file.Name()) |
| 117 | |
| 118 | ctx := t.Context() |
| 119 | |
| 120 | token := &protocol.ClientToken{ |
| 121 | Token: []byte("test"), |
| 122 | } |
| 123 | cl := &protocol.Node{ |
| 124 | Id: chord.Random(), |
| 125 | } |
| 126 | |
| 127 | hostname := "custom.domain.com" |
| 128 | |
| 129 | der, cert, key := makeCertificate(as, logger, cl, token, nil) |
| 130 | privKey, err := pki.UnmarshalPrivateKey([]byte(key)) |
| 131 | as.NoError(err) |
| 132 | |
| 133 | cfg := &Config{ |
| 134 | path: file.Name(), |
| 135 | router: skipmap.NewString[route](), |
| 136 | Apex: testApex, |
| 137 | Certificate: cert, |
| 138 | PrivKey: key, |
| 139 | Tunnels: []Tunnel{ |
| 140 | { |
| 141 | Target: "tcp://127.0.0.1:2345", |
| 142 | }, |
| 143 | }, |
| 144 | } |
| 145 | as.NoError(cfg.validate()) |
| 146 | |
| 147 | m := func(s *mocks.TunnelService, t1 *mocks.MemoryTransport, publishCall *mock.Call) { |
| 148 | s.On("AcmeValidate", mock.Anything, mock.MatchedBy(func(req *protocol.ValidateRequest) bool { |
| 149 | return powValidateFunc(hostname, privKey)(req.GetProof(), req.GetHostname()) |
| 150 | })).Return(&protocol.ValidateResponse{ |
| 151 | Apex: testApex, |
| 152 | }, nil) |
| 153 | |
| 154 | defaultNoHostnames(s) |
| 155 | transportHelper(t1, der) |
| 156 | } |
| 157 | |
| 158 | listenCfg := &net.ListenConfig{} |
| 159 | sListener, err := listenCfg.Listen(ctx, "tcp", "127.0.0.1:0") |
| 160 | as.NoError(err) |
| 161 | defer sListener.Close() |
| 162 | |
| 163 | client, _, assertion := setupClient(t, as, ctx, logger, nil, cfg, nil, m, false, 1) |
| 164 | defer assertion() |
| 165 | defer client.Close() |
| 166 | |
| 167 | client.ServerListener = sListener |
| 168 | |
| 169 | client.Start(ctx) |
| 170 | |
| 171 | resp, err := client.RequestAcmeValidation(ctx, hostname) |
| 172 | as.NoError(err) |
| 173 | as.NotNil(resp) |
| 174 | as.EqualValues(testApex, resp.GetApex()) |
| 175 | |
| 176 | c := &http.Client{ |
| 177 | Timeout: acme.HashcashExpires, |
| 178 | } |
| 179 | httpReq, err := http.NewRequest(http.MethodGet, fmt.Sprintf("http://%s/api/validate/%s", sListener.Addr().String(), hostname), nil) |
| 180 | as.NoError(err) |
| 181 | httpResp, err := c.Do(httpReq) |
| 182 | as.NoError(err) |
| 183 | defer httpResp.Body.Close() |
| 184 | |
| 185 | body, err := io.ReadAll(httpResp.Body) |
| 186 | as.NoError(err) |
| 187 | as.Contains(string(body), resp.GetApex()) |
| 188 | } |
| 189 | |
| 190 | func powValidateFunc(expectHostname string, privKey ed25519.PrivateKey) func(reqProof *protocol.ProofOfWork, reqHostname string) bool { |
| 191 | return func(reqProof *protocol.ProofOfWork, reqHostname string) bool { |
| 192 | match := reqHostname == expectHostname |
| 193 | if !match { |
| 194 | return false |
| 195 | } |
| 196 | d, err := pow.VerifySolution(reqProof, pow.Parameters{ |
| 197 | Difficulty: acme.HashcashDifficulty, |
| 198 | Expires: acme.HashcashExpires, |
| 199 | GetSubject: func(pubKey ed25519.PublicKey) string { |
| 200 | return reqHostname |
| 201 | }, |
| 202 | }) |
| 203 | if err != nil { |
| 204 | return false |
| 205 | } |
| 206 | return bytes.Equal(d.PubKey, privKey.Public().(ed25519.PublicKey)) |
| 207 | } |
| 208 | } |