aboutsummaryrefslogtreecommitdiffstats
path: root/dec.go
blob: 357f97934d45c9d7e4964950e08a47cdce50d980 (plain) (blame)
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
127
128
129
package main

import (
	"context"
	"errors"
	"flag"
	"fmt"
	"io"
	"os"
	"strings"

	"github.com/google/subcommands"
)

type decCmd struct {
	password string
	force    bool

	passwordIn func() (string, error)
	stdin      io.Reader
	stdout     io.Writer
}

func (*decCmd) Name() string     { return "dec" }
func (*decCmd) Synopsis() string { return "decrypt" }
func (*decCmd) Usage() string {
	return `usage: sym dec [OPTION]... [FILE]...
Decrypt files, or stdin if no files are provided.

-p is required when reading from stdin.

For example,
  sym dec my-encrypted-file.txt.enc
would decrypt my-encrypted-file.txt.enc and write the result to
my-encrypted-file.txt. If a filename does not end with .enc, the name
will be appended with a .dec extension.

`
}

func (c *decCmd) SetFlags(fs *flag.FlagSet) {
	fs.StringVar(&c.password, "p", "", "use the specified password; if not provided, dec will prompt for a password")
	fs.BoolVar(&c.force, "f", false, "overwrite output files even if they already exist")
}

func (c *decCmd) decrypt(w io.Writer, r io.Reader, password string) error {
	_, err := io.Copy(w, newDecryptingReader(r, password))
	return err
}

func (c *decCmd) decryptFile(fileName string, password string) (err error) {
	var outFileName string
	if name, ok := strings.CutSuffix(fileName, ".enc"); ok {
		outFileName = name
	} else {
		outFileName = fileName + ".dec"
	}
	fIn, err := os.Open(fileName)
	if err != nil {
		return err
	}
	defer fIn.Close()
	fileOpts := os.O_CREATE | os.O_WRONLY
	if c.force {
		fileOpts |= os.O_TRUNC
	} else {
		fileOpts |= os.O_EXCL
	}
	fOut, err := os.OpenFile(outFileName, fileOpts, 0644)
	if err != nil {
		if errors.Is(err, os.ErrExist) {
			return fmt.Errorf("output file %q exists (use -f to overwrite)", outFileName)
		}
		return err
	}
	defer func() {
		fOut.Close()
		if err != nil {
			os.Remove(fOut.Name())
		}
	}()
	if err := c.decrypt(fOut, fIn, password); err != nil {
		return fmt.Errorf("decrypt %q: %s", fileName, err)
	}
	return fOut.Close()
}

func (c *decCmd) readPassword() (string, error) {
	fmt.Fprint(os.Stderr, "Enter password: ")
	pw, err := c.passwordIn()
	fmt.Fprintln(os.Stderr)
	return pw, err
}

func (c *decCmd) run(args ...string) error {
	if len(args) == 0 && c.password == "" {
		return usageErr("-p is required when reading from stdin")
	}
	var password string
	if c.password != "" {
		password = c.password
	} else {
		var err error
		password, err = c.readPassword()
		if err != nil {
			return err
		}
	}
	if len(args) == 0 {
		return c.decrypt(c.stdout, c.stdin, password)
	}
	for _, fileName := range args {
		if err := c.decryptFile(fileName, password); err != nil {
			return err
		}
	}
	return nil
}

func (c *decCmd) Execute(ctx context.Context, f *flag.FlagSet, _ ...any) subcommands.ExitStatus {
	if err := c.run(f.Args()...); err != nil {
		fmt.Fprintf(os.Stderr, "sym: %s\n", err)
		if errors.Is(err, errUsage) {
			return subcommands.ExitUsageError
		}
		return subcommands.ExitFailure
	}
	return subcommands.ExitSuccess
}