Skip to content
File

Blob: tun/client/acme_test.go

go209 lines
1package client
2 
3import (
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 
26func 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 
110func 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 
190func 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}