Repository navigation
Expand file tree
/
Copy pathupdate_test.go
More file actions
126 lines (118 loc) · 3.72 KB
/
Copy pathupdate_test.go
File metadata and controls
126 lines (118 loc) · 3.72 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
package main
import (
"crypto/ed25519"
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"encoding/json"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"testing"
"time"
"wanctl/internal/config"
wanrelease "wanctl/internal/release"
)
// signedUpdateServer serves one release the way a relay's /dl mirror or a
// GitHub release page does: manifest, detached signature, artifact. The OS,
// arch and version are parameters because the auto-updater downloads for the
// platform it is actually running on, so its tests cannot use a fixture that is
// permanently linux/amd64.
func signedUpdateServer(t *testing.T, payload []byte, goos, goarch, version string, mutate func(path string, body []byte) []byte) *httptest.Server {
t.Helper()
pub, priv, err := ed25519.GenerateKey(rand.Reader)
if err != nil {
t.Fatal(err)
}
wanrelease.TrustedPublicKeys = base64.StdEncoding.EncodeToString(pub)
t.Cleanup(func() { wanrelease.TrustedPublicKeys = "" })
h := sha256.Sum256(payload)
artifact := wanrelease.ArtifactName(goos, goarch)
manifest, err := json.Marshal(wanrelease.Manifest{
Schema: 1, Version: version, PublishedAt: time.Now().UTC(),
Artifacts: []wanrelease.Artifact{{OS: goos, Arch: goarch, Name: artifact, Size: int64(len(payload)), SHA256: hex.EncodeToString(h[:])}},
})
if err != nil {
t.Fatal(err)
}
signature := ed25519.Sign(priv, manifest)
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
var body []byte
switch req.URL.Path {
case "/dl/" + wanrelease.ManifestName:
body = manifest
case "/dl/" + wanrelease.SignatureName:
body = signature
case "/dl/" + artifact:
body = payload
default:
http.NotFound(w, req)
return
}
if mutate != nil {
body = mutate(req.URL.Path, body)
}
w.Write(body)
}))
}
func TestDownloadSignedUpdateVerifiesBeforeReturning(t *testing.T) {
srv := signedUpdateServer(t, []byte("signed binary"), "linux", "amd64", "v2.0.0", nil)
defer srv.Close()
dir := t.TempDir()
path, version, err := downloadSignedUpdate(t.Context(), srv.URL+"/dl", dir, "linux", "amd64", "v1.0.0")
if err != nil {
t.Fatal(err)
}
defer os.Remove(path)
if version != "v2.0.0" {
t.Fatalf("version = %q", version)
}
got, err := os.ReadFile(path)
if err != nil || string(got) != "signed binary" {
t.Fatalf("downloaded = %q, %v", got, err)
}
if filepath.Dir(path) != dir {
t.Fatalf("temporary artifact is outside destination directory: %s", path)
}
}
func TestDownloadSignedUpdateRejectsTamperedArtifact(t *testing.T) {
srv := signedUpdateServer(t, []byte("signed binary"), "linux", "amd64", "v2.0.0", func(path string, body []byte) []byte {
if path == "/dl/wanctl-linux-amd64" {
return []byte("attacker binary")
}
return body
})
defer srv.Close()
dir := t.TempDir()
if _, _, err := downloadSignedUpdate(t.Context(), srv.URL+"/dl", dir, "linux", "amd64", "v1.0.0"); err == nil {
t.Fatal("tampered artifact accepted")
}
entries, err := os.ReadDir(dir)
if err != nil {
t.Fatal(err)
}
if len(entries) != 0 {
t.Fatalf("failed update left files behind: %v", entries)
}
}
func TestPlanUpdateRestartPreservesSupervisorOwnership(t *testing.T) {
dir := t.TempDir()
t.Setenv("WANCTL_CONFIG_DIR", dir)
const pid = 4242
if err := config.WriteManagedPID(pid); err != nil {
t.Fatal(err)
}
plan := planUpdateRestartWithLiveness(false, pid, true)
if plan.stopDetached || plan.restartDetached || plan.restartManagedPID != pid {
t.Fatalf("managed plan = %+v", plan)
}
if err := config.WriteManagedPID(pid + 1); err != nil {
t.Fatal(err)
}
plan = planUpdateRestartWithLiveness(false, pid, true)
if !plan.stopDetached || !plan.restartDetached || plan.restartManagedPID != 0 {
t.Fatalf("stale managed marker plan = %+v", plan)
}
}