Skip to content

Commit 196b8f4

Browse files
committed
more config validation, refactoring
1 parent a549bfa commit 196b8f4

2 files changed

Lines changed: 146 additions & 102 deletions

File tree

extract.go

Lines changed: 77 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,77 @@
1+
//go:debug tarinsecurepath=0
2+
3+
package main
4+
5+
import (
6+
"archive/tar"
7+
"bytes"
8+
"compress/gzip"
9+
"fmt"
10+
"io"
11+
"os"
12+
"path/filepath"
13+
"strings"
14+
)
15+
16+
func extractTarGz(tarGzData []byte, destination string, stripComponents int) error {
17+
buf := bytes.NewBuffer(tarGzData)
18+
gzipReader, err := gzip.NewReader(buf)
19+
if err != nil {
20+
return err
21+
}
22+
defer gzipReader.Close()
23+
tarReader := tar.NewReader(gzipReader)
24+
25+
for {
26+
header, err := tarReader.Next()
27+
if err == io.EOF {
28+
break
29+
}
30+
if err != nil {
31+
return err
32+
}
33+
// Skip pax_global_header entries
34+
if header.Name == "pax_global_header" {
35+
continue
36+
}
37+
38+
// Calculate the target path by stripping components
39+
target := header.Name
40+
if stripComponents > 0 {
41+
components := strings.SplitN(target, string(filepath.Separator), stripComponents+1)
42+
if len(components) > stripComponents {
43+
target = strings.Join(components[stripComponents:], string(filepath.Separator))
44+
} else {
45+
target = ""
46+
}
47+
}
48+
49+
// Get the full path for the file
50+
target = filepath.Join(destination, target)
51+
52+
switch header.Typeflag {
53+
case tar.TypeDir:
54+
// Create directory if it doesn't exist
55+
if err := os.MkdirAll(target, os.ModePerm); err != nil {
56+
return err
57+
}
58+
59+
case tar.TypeReg:
60+
// Create file
61+
file, err := os.Create(target)
62+
if err != nil {
63+
return err
64+
}
65+
defer file.Close()
66+
67+
if _, err := io.Copy(file, tarReader); err != nil {
68+
return err
69+
}
70+
71+
default:
72+
return fmt.Errorf("unsupported file type: %v in %v", header.Typeflag, header.Name)
73+
}
74+
}
75+
76+
return nil
77+
}

updater.go

Lines changed: 69 additions & 102 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,6 @@
1-
//go:debug tarinsecurepath=0
2-
31
package main
42

53
import (
6-
"archive/tar"
7-
"bytes"
8-
"compress/gzip"
94
"context"
105
"encoding/json"
116
"flag"
@@ -32,7 +27,6 @@ type Target struct {
3227
}
3328

3429
type Config struct {
35-
MetadataDir string `json:"metadata_dir"`
3630
DeployDir string `json:"deploy_dir"`
3731
Targets []*Target `json:"targets"`
3832
PublicSigningKey string `json:"public_signing_key"`
@@ -55,19 +49,14 @@ func main() {
5549
log.Fatal(err)
5650
}
5751

58-
err = validateConfig(&config)
59-
if err != nil {
60-
log.Fatal(err)
61-
}
62-
63-
err = os.MkdirAll(config.MetadataDir, 0750)
52+
err = config.Validate()
6453
if err != nil {
6554
log.Fatal(err)
6655
}
6756

6857
githubAPIToken, ok := os.LookupEnv("GITHUB_API_TOKEN")
69-
if !ok {
70-
log.Fatal("GITHUB_API_TOKEN environment variable must be set")
58+
if !ok || githubAPIToken == "" {
59+
log.Fatal("GITHUB_API_TOKEN environment variable must be set and non-empty")
7160
}
7261
client := github.NewClient(nil).WithAuthToken(githubAPIToken)
7362

@@ -83,14 +72,14 @@ func main() {
8372
}
8473

8574
releaseID := strconv.FormatInt(*release.ID, 10)
86-
lastReleaseFile := filepath.Join(config.MetadataDir, fmt.Sprintf("%v_last_release_id", target.Name))
8775
missingLastRelease := false
88-
lastReleaseID, err := os.ReadFile(lastReleaseFile)
76+
lastReleaseID, err := getLastReleaseID(config.DeployDir, target.Name)
8977
if err != nil {
9078
if os.IsNotExist(err) {
9179
missingLastRelease = true
9280
} else {
93-
log.Fatal(err)
81+
log.Printf("%s: error getting last release ID: %v", target.Name, err)
82+
continue
9483
}
9584
}
9685
if !missingLastRelease && string(lastReleaseID) == releaseID {
@@ -115,7 +104,7 @@ func main() {
115104
log.Printf("%s: skipping signature verification!", target.Name)
116105
}
117106

118-
err = deploy(config.DeployDir, target.Name, releaseID, tarGzBytes, lastReleaseFile, string(lastReleaseID))
107+
err = deployRelease(config.DeployDir, target.Name, releaseID, lastReleaseID, tarGzBytes)
119108
if err != nil {
120109
log.Printf("%s update failed: %v", target.Name, err)
121110
continue
@@ -126,25 +115,67 @@ func main() {
126115
}
127116
}
128117

129-
func validateConfig(config *Config) error {
130-
if config.MetadataDir == "" {
131-
return fmt.Errorf("metadata directory must be set")
132-
}
133-
if config.DeployDir == "" {
118+
func (c *Config) Validate() error {
119+
if c.DeployDir == "" {
134120
return fmt.Errorf("deploy directory must be set")
135121
}
136-
if config.UpdateInterval <= 0 {
122+
if c.UpdateInterval <= 0 {
137123
return fmt.Errorf("update interval must be >0")
138124
}
139-
if !config.UnsafeSkipSignatureVerification && config.PublicSigningKey == "" {
125+
if !c.UnsafeSkipSignatureVerification && c.PublicSigningKey == "" {
140126
return fmt.Errorf("public signing key must be set if signature verification is enabled")
141127
}
142-
if len(config.Targets) == 0 {
128+
if len(c.Targets) == 0 {
143129
return fmt.Errorf("at least one target must be set")
144130
}
131+
targetNames := make(map[string]bool)
132+
for i, target := range c.Targets {
133+
if target.Name == "" {
134+
return fmt.Errorf("name for target %d must be set", i)
135+
}
136+
if target.Owner == "" {
137+
return fmt.Errorf("owner for target %d must be set", i)
138+
}
139+
if target.Repo == "" {
140+
return fmt.Errorf("repo for target %d must be set", i)
141+
}
142+
if targetNames[target.Name] {
143+
return fmt.Errorf("target %d has duplicate name", i)
144+
}
145+
targetNames[target.Name] = true
146+
}
145147
return nil
146148
}
147149

150+
func getReleaseDir(deployDir, targetName, releaseID string) string {
151+
return filepath.Join(deployDir, targetName+"-"+releaseID)
152+
}
153+
154+
func getReleaseSymlink(deployDir, targetName string) string {
155+
return filepath.Join(deployDir, targetName)
156+
}
157+
158+
func getLastReleaseID(deployDir, targetName string) (string, error) {
159+
lastReleaseSymlink := getReleaseSymlink(deployDir, targetName)
160+
fi, err := os.Lstat(lastReleaseSymlink)
161+
if err != nil {
162+
return "", err
163+
}
164+
if fi.Mode()&os.ModeSymlink == 0 {
165+
return "", fmt.Errorf("%s is not a symlink", lastReleaseSymlink)
166+
}
167+
168+
lastReleaseDir, err := os.Readlink(lastReleaseSymlink)
169+
if err != nil {
170+
return "", err
171+
}
172+
split := strings.Split(filepath.Base(lastReleaseDir), "-")
173+
if len(split) != 2 {
174+
return "", fmt.Errorf("invalid last release directory name: %s", lastReleaseDir)
175+
}
176+
return split[1], nil
177+
}
178+
148179
func downloadReleaseAssets(target *Target, release *github.RepositoryRelease) (tarGzBytes, sigBytes []byte, err error) {
149180
if len(release.Assets) < 2 {
150181
err = fmt.Errorf("release needs at least 2 assets (have %v)", len(release.Assets))
@@ -153,8 +184,8 @@ func downloadReleaseAssets(target *Target, release *github.RepositoryRelease) (t
153184

154185
const tarGzRegexFmt = `^%s-[\w.]+\.tar\.gz$`
155186
const sigRegexFmt = `^%s-[\w.]+\.minisig$`
156-
tarGzRegex := regexp.MustCompile(fmt.Sprintf(tarGzRegexFmt, target.Name))
157-
sigRegex := regexp.MustCompile(fmt.Sprintf(sigRegexFmt, target.Name))
187+
tarGzRegex := regexp.MustCompile(fmt.Sprintf(tarGzRegexFmt, target.Repo))
188+
sigRegex := regexp.MustCompile(fmt.Sprintf(sigRegexFmt, target.Repo))
158189

159190
if !(tarGzRegex.MatchString(*release.Assets[0].Name)) {
160191
err = fmt.Errorf("first asset doesn't have expected name (%v)", *release.Assets[0].Name)
@@ -222,87 +253,23 @@ func verifySignature(publicSigningKey string, tarGzBytes, sigBytes []byte) (bool
222253
return pk.Verify(tarGzBytes, sig)
223254
}
224255

225-
func deploy(deployDir, targetName, releaseID string, tarGzBytes []byte, lastReleaseFile, lastReleaseID string) error {
226-
extractDir := filepath.Join(deployDir, targetName) + "-" + releaseID
227-
if err := os.Mkdir(extractDir, 0755); err != nil {
256+
func deployRelease(deployDir, targetName, releaseID, lastReleaseID string, tarGzBytes []byte) error {
257+
releaseDir := getReleaseDir(deployDir, targetName, releaseID)
258+
if err := os.Mkdir(releaseDir, 0755); err != nil {
228259
return err
229260
}
230-
if err := extractTarGz(tarGzBytes, extractDir, 1); err != nil {
231-
return err
232-
}
233-
if err := os.Symlink(extractDir, extractDir+".tmp"); err != nil {
234-
return err
235-
}
236-
if err := os.Rename(extractDir+".tmp", filepath.Join(deployDir, targetName)); err != nil {
237-
return err
238-
}
239-
if err := os.WriteFile(lastReleaseFile, []byte(releaseID), 0640); err != nil {
261+
if err := extractTarGz(tarGzBytes, releaseDir, 1); err != nil {
240262
return err
241263
}
242264

243-
// clean up old release dir
244-
return os.RemoveAll(filepath.Join(deployDir, targetName) + "-" + lastReleaseID)
245-
}
246-
247-
func extractTarGz(tarGzData []byte, destination string, stripComponents int) error {
248-
buf := bytes.NewBuffer(tarGzData)
249-
gzipReader, err := gzip.NewReader(buf)
250-
if err != nil {
265+
releaseSymlink := getReleaseSymlink(deployDir, targetName)
266+
if err := os.Symlink(releaseDir, releaseSymlink+".tmp"); err != nil {
251267
return err
252268
}
253-
defer gzipReader.Close()
254-
tarReader := tar.NewReader(gzipReader)
255-
256-
for {
257-
header, err := tarReader.Next()
258-
if err == io.EOF {
259-
break
260-
}
261-
if err != nil {
262-
return err
263-
}
264-
// Skip pax_global_header entries
265-
if header.Name == "pax_global_header" {
266-
continue
267-
}
268-
269-
// Calculate the target path by stripping components
270-
target := header.Name
271-
if stripComponents > 0 {
272-
components := strings.SplitN(target, string(filepath.Separator), stripComponents+1)
273-
if len(components) > stripComponents {
274-
target = strings.Join(components[stripComponents:], string(filepath.Separator))
275-
} else {
276-
target = ""
277-
}
278-
}
279-
280-
// Get the full path for the file
281-
target = filepath.Join(destination, target)
282-
283-
switch header.Typeflag {
284-
case tar.TypeDir:
285-
// Create directory if it doesn't exist
286-
if err := os.MkdirAll(target, os.ModePerm); err != nil {
287-
return err
288-
}
289-
290-
case tar.TypeReg:
291-
// Create file
292-
file, err := os.Create(target)
293-
if err != nil {
294-
return err
295-
}
296-
defer file.Close()
297-
298-
if _, err := io.Copy(file, tarReader); err != nil {
299-
return err
300-
}
301-
302-
default:
303-
return fmt.Errorf("unsupported file type: %v in %v", header.Typeflag, header.Name)
304-
}
269+
if err := os.Rename(releaseSymlink+".tmp", releaseSymlink); err != nil {
270+
return err
305271
}
306272

307-
return nil
273+
// clean up last release dir
274+
return os.RemoveAll(getReleaseDir(deployDir, targetName, lastReleaseID))
308275
}

0 commit comments

Comments
 (0)