summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--internal/api/api.go9
-rw-r--r--internal/api/errorcode_string.go9
-rw-r--r--internal/cryptoutil/cryptoutil.go26
-rw-r--r--internal/cryptoutil/cryptoutil_test.go20
-rw-r--r--roseh.moe.go138
-rw-r--r--tools/notes/notes.go108
6 files changed, 199 insertions, 111 deletions
diff --git a/internal/api/api.go b/internal/api/api.go
index 7c8472a..b925a7a 100644
--- a/internal/api/api.go
+++ b/internal/api/api.go
@@ -11,6 +11,7 @@ const (
Unauthenticated // missing credentials
PermissionDenied // invalid credentials, or login required
BadRequest // incorrect request format (try gob)
+ InvalidArgument // one or more arguments was invalid
NotFound // requested resource was not found
Internal // server encountered an unexpected error
)
@@ -35,12 +36,14 @@ type ListNotesResponse struct {
Notes []string
}
-type CreateNoteRequestStream struct {
- Chunk []byte
+type CreateNoteRequest struct {
+ ContinuationToken []byte
+ Chunk []byte
}
type CreateNoteResponse struct {
- Name string
+ Name string
+ ContinuationToken []byte
}
type ReadNoteRequest struct {
diff --git a/internal/api/errorcode_string.go b/internal/api/errorcode_string.go
index 2b4a294..5ff2ff0 100644
--- a/internal/api/errorcode_string.go
+++ b/internal/api/errorcode_string.go
@@ -12,13 +12,14 @@ func _() {
_ = x[Unauthenticated-1]
_ = x[PermissionDenied-2]
_ = x[BadRequest-3]
- _ = x[NotFound-4]
- _ = x[Internal-5]
+ _ = x[InvalidArgument-4]
+ _ = x[NotFound-5]
+ _ = x[Internal-6]
}
-const _ErrorCode_name = "OkUnauthenticatedPermissionDeniedBadRequestNotFoundInternal"
+const _ErrorCode_name = "OkUnauthenticatedPermissionDeniedBadRequestInvalidArgumentNotFoundInternal"
-var _ErrorCode_index = [...]uint8{0, 2, 17, 33, 43, 51, 59}
+var _ErrorCode_index = [...]uint8{0, 2, 17, 33, 43, 58, 66, 74}
func (i ErrorCode) String() string {
if i < 0 || i >= ErrorCode(len(_ErrorCode_index)-1) {
diff --git a/internal/cryptoutil/cryptoutil.go b/internal/cryptoutil/cryptoutil.go
index b269803..0429b02 100644
--- a/internal/cryptoutil/cryptoutil.go
+++ b/internal/cryptoutil/cryptoutil.go
@@ -9,8 +9,10 @@ import (
"crypto/rand"
"crypto/sha512"
"crypto/subtle"
+ "encoding/binary"
"encoding/hex"
"errors"
+ "io"
)
const (
@@ -101,23 +103,31 @@ func (h PasswordHash) CheckPassword(password string) (RawKey, error) {
return key, nil
}
-// Sign generates an HMAC-SHA512 signature and appends it to msg.
-func (k HMACKey) Sign(msg []byte) SignedMessage {
+func (k HMACKey) mac(out, msg []byte, info string) []byte {
mac := hmac.New(sha512.New, k)
+ buf := make([]byte, 0, binary.MaxVarintLen64)
+ mac.Write(binary.AppendUvarint(buf, uint64(len(info))))
+ io.WriteString(mac, info)
mac.Write(msg)
- return mac.Sum(msg)
+ return mac.Sum(out)
+}
+
+// Sign generates an HMAC-SHA512 signature and appends it to msg. The info and
+// additionalData are also authenticated, but are not included in the returned
+// signed message.
+func (k HMACKey) Sign(msg []byte, info string) SignedMessage {
+ return k.mac(msg, msg, info)
}
// Verify checks whether the given message has a valid signature, and returns
-// the raw message if it does.
-func (k HMACKey) Verify(msg SignedMessage) ([]byte, bool) {
+// the raw message if it does. The additionalData must match the data passed
+// to Sign.
+func (k HMACKey) Verify(msg SignedMessage, info string) ([]byte, bool) {
if len(msg) < certSize {
return nil, false
}
msg, sig := msg[:len(msg)-certSize], msg[len(msg)-certSize:]
- mac := hmac.New(sha512.New, k)
- mac.Write(msg)
- if !hmac.Equal(sig, mac.Sum(nil)) {
+ if !hmac.Equal(sig, k.mac(nil, msg, info)) {
return nil, false
}
return msg, true
diff --git a/internal/cryptoutil/cryptoutil_test.go b/internal/cryptoutil/cryptoutil_test.go
index 1670eff..0e0560c 100644
--- a/internal/cryptoutil/cryptoutil_test.go
+++ b/internal/cryptoutil/cryptoutil_test.go
@@ -31,8 +31,8 @@ func TestPassword(t *testing.T) {
func TestSignature(t *testing.T) {
key := HMACKey(mustHex(t, "669e06ec457778b9a8133edb0a87ea82c6b141ffbbc63c038da96258175eb35c"))
msg := []byte("test message")
- signedMsg := key.Sign(msg)
- got, ok := key.Verify(signedMsg)
+ signedMsg := key.Sign(msg, "info")
+ got, ok := key.Verify(signedMsg, "info")
if !ok {
t.Fatalf("Verify(%x) rejected the message", signedMsg)
}
@@ -64,14 +64,14 @@ func TestEncrypt(t *testing.T) {
}} {
t.Run(tc.desc, func(t *testing.T) {
encryptedMsg := new(bytes.Buffer)
- w := key.NewWriter(encryptedMsg, nil)
+ w := key.NewWriter(encryptedMsg, []byte("additional data"))
if _, err := w.Write(tc.msg); err != nil {
t.Fatalf("EncryptingWriter.Write(%q) failed: %s", tc.msg, err)
}
if err := w.Close(); err != nil {
t.Fatalf("EncryptingWriter.Close() failed: %s", err)
}
- got, err := io.ReadAll(key.NewReader(bytes.NewReader(encryptedMsg.Bytes()), nil))
+ got, err := io.ReadAll(key.NewReader(bytes.NewReader(encryptedMsg.Bytes()), []byte("additional data")))
if err != nil {
t.Fatalf("DecryptingReader.Read(%x) failed: %s", encryptedMsg, err)
}
@@ -98,14 +98,14 @@ func TestDeriveKey(t *testing.T) {
}
msg := []byte("test message")
encryptedMsg := new(bytes.Buffer)
- w := key.NewWriter(encryptedMsg, nil)
+ w := key.NewWriter(encryptedMsg, []byte("additional data"))
if _, err := w.Write(msg); err != nil {
t.Fatalf("EncryptingWriter.Write(%q) failed: %s", msg, err)
}
if err := w.Close(); err != nil {
t.Fatalf("EncryptingWriter.Close() failed: %s", err)
}
- got, err := io.ReadAll(key.NewReader(encryptedMsg, nil))
+ got, err := io.ReadAll(key.NewReader(encryptedMsg, []byte("additional data")))
if err != nil {
t.Fatalf("DecryptingReader.Read(%x) failed: %s", encryptedMsg, err)
}
@@ -121,8 +121,9 @@ func BenchmarkEncrypt(b *testing.B) {
msg[i] = byte(i)
}
encryptedMsg := bytes.NewBuffer(make([]byte, 0, 35+len(msg)+16*(len(msg)+segmentSize-1)/segmentSize /* ?? */))
+ additionalData := []byte("additional data")
for b.Loop() {
- w := key.NewWriter(encryptedMsg, nil)
+ w := key.NewWriter(encryptedMsg, additionalData)
if _, err := w.Write(msg); err != nil {
b.Fatal(err)
}
@@ -140,7 +141,8 @@ func BenchmarkDecrypt(b *testing.B) {
msg[i] = byte(i)
}
encryptedMsg := new(bytes.Buffer)
- w := key.NewWriter(encryptedMsg, nil)
+ additionalData := []byte("additional data")
+ w := key.NewWriter(encryptedMsg, additionalData)
if _, err := w.Write(msg); err != nil {
b.Fatal(err)
}
@@ -149,7 +151,7 @@ func BenchmarkDecrypt(b *testing.B) {
}
decryptedMsg := make([]byte, len(msg))
for b.Loop() {
- r := key.NewReader(bytes.NewReader(encryptedMsg.Bytes()), nil)
+ r := key.NewReader(bytes.NewReader(encryptedMsg.Bytes()), additionalData)
if _, err := io.ReadFull(r, decryptedMsg); err != nil {
b.Fatal(err)
}
diff --git a/roseh.moe.go b/roseh.moe.go
index f9c5478..76b4dc1 100644
--- a/roseh.moe.go
+++ b/roseh.moe.go
@@ -133,12 +133,16 @@ func favicon(w http.ResponseWriter, r *http.Request) {
serveStaticFile(w, r, "static/favicon.ico")
}
+type authToken struct {
+ Issued time.Time
+}
+
func makeToken() (string, error) {
- nowBytes, err := time.Now().MarshalBinary()
- if err != nil {
+ buf := new(bytes.Buffer)
+ if err := gob.NewEncoder(buf).Encode(authToken{Issued: time.Now()}); err != nil {
return "", err
}
- return base64.RawStdEncoding.EncodeToString(secretKey.Sign(nowBytes)), nil
+ return base64.RawStdEncoding.EncodeToString(secretKey.Sign(buf.Bytes(), "auth")), nil
}
func checkToken(token string) bool {
@@ -146,15 +150,15 @@ func checkToken(token string) bool {
if err != nil {
return false
}
- msg, ok := secretKey.Verify(authCookie)
+ msg, ok := secretKey.Verify(authCookie, "auth")
if !ok {
return false
}
- var t time.Time
- if err := t.UnmarshalBinary(msg); err != nil {
+ var t authToken
+ if err := gob.NewDecoder(bytes.NewReader(msg)).Decode(&t); err != nil {
return false
}
- return time.Since(t) < cookieExpiration
+ return time.Since(t.Issued) < cookieExpiration
}
const cookieExpiration = 180 * 24 * time.Hour
@@ -402,13 +406,84 @@ func listNotes(req *api.ListNotesRequest) (*api.ListNotesResponse, error) {
return &api.ListNotesResponse{Notes: names}, nil
}
+type countingWriter struct {
+ w io.Writer
+ count int64
+}
+
+func (w *countingWriter) Write(buf []byte) (int, error) {
+ n, err := w.w.Write(buf)
+ w.count += int64(n)
+ return n, err
+}
+
+type continuationToken struct {
+ Name string
+ Size int64
+}
+
+func appendNote(req *api.CreateNoteRequest, key cryptoutil.EncryptionKey) (*api.CreateNoteResponse, error) {
+ msg, ok := secretKey.Verify(req.ContinuationToken, "continuation-token")
+ if !ok {
+ return nil, fmt.Errorf("%w: bad continuation token", api.InvalidArgument)
+ }
+ var token continuationToken
+ if err := gob.NewDecoder(bytes.NewReader(msg)).Decode(&token); err != nil {
+ return nil, fmt.Errorf("failed to decode continuation token: %s", err)
+ }
+ var err error
+ oldNote, err := os.Open(*notepadDir + "/notes/" + token.Name)
+ if err != nil {
+ return nil, err
+ }
+ defer oldNote.Close()
+ stat, err := oldNote.Stat()
+ if err != nil {
+ return nil, err
+ }
+ if token.Size != stat.Size() {
+ return nil, fmt.Errorf("%w: continuation token expired", api.InvalidArgument)
+ }
+ newNote, err := os.CreateTemp(*notepadDir+"/notes", token.Name)
+ if err != nil {
+ return nil, err
+ }
+ defer newNote.Close()
+ additionalData := []byte("notes/" + token.Name)
+ cw := &countingWriter{w: newNote}
+ encryptingWriter := key.NewWriter(cw, additionalData)
+ if _, err := io.Copy(encryptingWriter, key.NewReader(oldNote, additionalData)); err != nil {
+ return nil, err
+ }
+ if _, err := encryptingWriter.Write(req.Chunk); err != nil {
+ return nil, err
+ }
+ if err := encryptingWriter.Close(); err != nil {
+ return nil, err
+ }
+ if err := newNote.Close(); err != nil {
+ return nil, err
+ }
+ if err := os.Rename(newNote.Name(), *notepadDir+"/notes/"+token.Name); err != nil {
+ return nil, err
+ }
+ buf := new(bytes.Buffer)
+ if err := gob.NewEncoder(buf).Encode(continuationToken{Name: token.Name, Size: cw.count}); err != nil {
+ return nil, err
+ }
+ return &api.CreateNoteResponse{
+ Name: token.Name,
+ ContinuationToken: secretKey.Sign(buf.Bytes(), "continuation-token"),
+ }, nil
+}
+
var (
//go:embed wordlist.txt
wordListString string
wordList = strings.Split(strings.TrimSuffix(wordListString, "\n"), "\n")
)
-func createNote(stream *gob.Decoder) (*api.CreateNoteResponse, error) {
+func createNote(req *api.CreateNoteRequest) (*api.CreateNoteResponse, error) {
encryptionKeyMu.Lock()
key := encryptionKey
encryptionKeyMu.Unlock()
@@ -417,6 +492,9 @@ func createNote(stream *gob.Decoder) (*api.CreateNoteResponse, error) {
}
var name string
var f *os.File
+ if len(req.ContinuationToken) > 0 {
+ return appendNote(req, key)
+ }
for n := 1; ; n++ {
buf := make([]byte, 2*n)
rand.Read(buf)
@@ -433,21 +511,13 @@ func createNote(stream *gob.Decoder) (*api.CreateNoteResponse, error) {
}
return nil, fmt.Errorf("create note: create note file: %s", err)
}
+ defer f.Close()
break
}
- defer f.Close()
- encryptingWriter := key.NewWriter(f, []byte("notes/"+name))
- for {
- req := new(api.CreateNoteRequestStream)
- if err := stream.Decode(req); err != nil {
- if errors.Is(err, io.EOF) {
- break
- }
- return nil, fmt.Errorf("create note: read input stream: %s", err)
- }
- if _, err := encryptingWriter.Write(req.Chunk); err != nil {
- return nil, fmt.Errorf("create note: write note to file: %s", err)
- }
+ cw := &countingWriter{w: f}
+ encryptingWriter := key.NewWriter(cw, []byte("notes/"+name))
+ if _, err := encryptingWriter.Write(req.Chunk); err != nil {
+ return nil, err
}
if err := encryptingWriter.Close(); err != nil {
return nil, err
@@ -455,7 +525,14 @@ func createNote(stream *gob.Decoder) (*api.CreateNoteResponse, error) {
if err := f.Close(); err != nil {
return nil, fmt.Errorf("create note: write file: %s", err)
}
- return &api.CreateNoteResponse{Name: name}, nil
+ token := new(bytes.Buffer)
+ if err := gob.NewEncoder(token).Encode(continuationToken{Name: name, Size: cw.count}); err != nil {
+ return nil, err
+ }
+ return &api.CreateNoteResponse{
+ Name: name,
+ ContinuationToken: secretKey.Sign(token.Bytes(), "continuation-token"),
+ }, nil
}
type readNoteResponseWriter struct {
@@ -521,21 +598,6 @@ func gobReqRespMiddleware[Request, Response any](next func(*Request) (*Response,
}
}
-func gobReqStreamMiddleware[Response any](next func(*gob.Decoder) (*Response, error)) http.HandlerFunc {
- return func(w http.ResponseWriter, r *http.Request) {
- w.Header().Set("Content-Type", "application/gob")
- if !tokenAuth(w, r) {
- return
- }
- resp, err := next(gob.NewDecoder(r.Body))
- if err != nil {
- gobError(w, err)
- return
- }
- writeGob(w, resp)
- }
-}
-
func gobRespStreamMiddleware[Request any](next func(*encoder, *Request) error) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/gob")
@@ -569,7 +631,7 @@ func main() {
http.HandleFunc("POST /notepad/autosave", autosave)
http.HandleFunc("POST /api/login", apiLogin)
http.HandleFunc("GET /api/list-notes", gobReqRespMiddleware(listNotes))
- http.HandleFunc("POST /api/create-note", gobReqStreamMiddleware(createNote))
+ http.HandleFunc("POST /api/create-note", gobReqRespMiddleware(createNote))
http.HandleFunc("GET /api/read-note", gobRespStreamMiddleware(readNote))
http.HandleFunc("GET /static/", static)
http.HandleFunc("GET /favicon.ico", favicon)
diff --git a/tools/notes/notes.go b/tools/notes/notes.go
index adc7e6a..2b03694 100644
--- a/tools/notes/notes.go
+++ b/tools/notes/notes.go
@@ -35,7 +35,7 @@ func readGobResp[Response any](req *http.Request) (*Response, error) {
return nil, fmt.Errorf("read gob response: decode response: %s", err)
}
if apiResp.Status != api.Ok {
- return nil, fmt.Errorf("read gob response: api error: %s", apiResp.Err)
+ return nil, fmt.Errorf("read gob response: api error: %s: %s", apiResp.Status, apiResp.Err)
}
resp, ok := apiResp.Ok.(*Response)
if !ok {
@@ -44,23 +44,19 @@ func readGobResp[Response any](req *http.Request) (*Response, error) {
return resp, nil
}
-func login(ctx context.Context) (string, error) {
+func login(ctx context.Context) (*api.LoginResponse, error) {
fmt.Print("Enter password:")
password, err := term.ReadPassword(int(os.Stdin.Fd()))
fmt.Println()
if err != nil {
- return "", err
+ return nil, err
}
req, err := http.NewRequestWithContext(ctx, "POST", *serverURL+"/api/login", nil)
if err != nil {
- return "", err
+ return nil, err
}
req.Header.Set("Roseh-Password", string(password))
- resp, err := readGobResp[api.LoginResponse](req)
- if err != nil {
- return "", err
- }
- return resp.Token, nil
+ return readGobResp[api.LoginResponse](req)
}
func loadToken(ctx context.Context) (string, error) {
@@ -71,8 +67,8 @@ func loadToken(ctx context.Context) (string, error) {
if err != nil {
return "", err
}
- os.WriteFile(tokenCache, []byte(token), 0600)
- return token, nil
+ os.WriteFile(tokenCache, []byte(token.Token), 0600)
+ return token.Token, nil
}
return "", err
}
@@ -96,6 +92,39 @@ func gobReq(ctx context.Context, method, url string, req any) (*http.Request, er
return httpReq, nil
}
+func gobReqResp[Response any](ctx context.Context, method, url string, req any) (*Response, error) {
+ httpReq, err := gobReq(ctx, method, url, req)
+ if err != nil {
+ return nil, err
+ }
+ return readGobResp[Response](httpReq)
+}
+
+func listNotes(ctx context.Context, req *api.ListNotesRequest) (*api.ListNotesResponse, error) {
+ return gobReqResp[api.ListNotesResponse](ctx, "GET", *serverURL+"/api/list-notes", req)
+}
+
+func createNote(ctx context.Context, req *api.CreateNoteRequest) (*api.CreateNoteResponse, error) {
+ return gobReqResp[api.CreateNoteResponse](ctx, "POST", *serverURL+"/api/create-note", req)
+}
+
+func readNote(ctx context.Context, req *api.ReadNoteRequest) (*gob.Decoder, func() error, error) {
+ httpReq, err := gobReq(ctx, "GET", *serverURL+"/api/read-note", req)
+ if err != nil {
+ return nil, nil, err
+ }
+ resp, err := http.DefaultClient.Do(httpReq)
+ if err != nil {
+ return nil, nil, err
+ }
+ if resp.StatusCode != http.StatusOK {
+ body, _ := io.ReadAll(resp.Body)
+ resp.Body.Close()
+ return nil, nil, fmt.Errorf("error status: %s\n%s", resp.Status, body)
+ }
+ return gob.NewDecoder(resp.Body), resp.Body.Close, nil
+}
+
type loginCommand struct{}
func (*loginCommand) Name() string {
@@ -117,7 +146,7 @@ func (*loginCommand) login(ctx context.Context) error {
if err != nil {
return err
}
- return os.WriteFile(tokenCache, []byte(token), 0600)
+ return os.WriteFile(tokenCache, []byte(token.Token), 0600)
}
func (c *loginCommand) Execute(ctx context.Context, _ *flag.FlagSet, _ ...any) subcommands.ExitStatus {
@@ -145,11 +174,7 @@ func (*listCommand) Usage() string {
func (*listCommand) SetFlags(*flag.FlagSet) {}
func (*listCommand) list(ctx context.Context) error {
- req, err := gobReq(ctx, "GET", *serverURL+"/api/list-notes", &api.ListNotesRequest{})
- if err != nil {
- return err
- }
- resp, err := readGobResp[api.ListNotesResponse](req)
+ resp, err := listNotes(ctx, &api.ListNotesRequest{})
if err != nil {
return err
}
@@ -184,13 +209,21 @@ func (*newCommand) Usage() string {
func (*newCommand) SetFlags(*flag.FlagSet) {}
type createNoteRequestStreamWriter struct {
- w *gob.Encoder
+ ctx context.Context
+ continuationToken []byte
+ name string
}
func (w *createNoteRequestStreamWriter) Write(buf []byte) (int, error) {
- if err := w.w.Encode(&api.CreateNoteRequestStream{Chunk: buf}); err != nil {
+ resp, err := createNote(w.ctx, &api.CreateNoteRequest{
+ ContinuationToken: w.continuationToken,
+ Chunk: buf,
+ })
+ if err != nil {
return 0, err
}
+ w.continuationToken = resp.ContinuationToken
+ w.name = resp.Name
return len(buf), nil
}
@@ -200,25 +233,11 @@ func (*newCommand) new(ctx context.Context, fileName string) error {
return err
}
defer f.Close()
- r, w := io.Pipe()
- req, err := http.NewRequestWithContext(ctx, "POST", *serverURL+"/api/create-note", r)
- if err != nil {
- return fmt.Errorf("request: %s", err)
- }
- token, err := loadToken(ctx)
- if err != nil {
- return fmt.Errorf("token: %s", err)
- }
- req.Header.Set("Roseh-Token", token)
- go func() {
- io.Copy(&createNoteRequestStreamWriter{gob.NewEncoder(w)}, f)
- w.Close()
- }()
- resp, err := readGobResp[api.CreateNoteResponse](req)
- if err != nil {
- return fmt.Errorf("req: %s", err)
+ streamWriter := &createNoteRequestStreamWriter{ctx: ctx}
+ if _, err := io.Copy(streamWriter, f); err != nil {
+ return err
}
- fmt.Println(resp.Name)
+ fmt.Println(streamWriter.name)
return nil
}
@@ -252,20 +271,11 @@ func (*readCommand) Usage() string {
func (*readCommand) SetFlags(*flag.FlagSet) {}
func (*readCommand) read(ctx context.Context, key string) error {
- req, err := gobReq(ctx, "GET", *serverURL+"/api/read-note", &api.ReadNoteRequest{Note: key})
+ decoder, close, err := readNote(ctx, &api.ReadNoteRequest{Note: key})
if err != nil {
return err
}
- resp, err := http.DefaultClient.Do(req)
- if err != nil {
- return fmt.Errorf("http: %s", err)
- }
- defer resp.Body.Close()
- if resp.StatusCode != http.StatusOK {
- body, _ := io.ReadAll(resp.Body)
- return fmt.Errorf("error status: %s\n%s", resp.Status, body)
- }
- decoder := gob.NewDecoder(resp.Body)
+ defer close()
for {
apiResp := new(api.Response)
if err := decoder.Decode(apiResp); err != nil {
@@ -275,7 +285,7 @@ func (*readCommand) read(ctx context.Context, key string) error {
return err
}
if apiResp.Status != api.Ok {
- return fmt.Errorf("api error: %s", apiResp.Err)
+ return fmt.Errorf("%s: %s", apiResp.Status, apiResp.Err)
}
resp, ok := apiResp.Ok.(*api.ReadNoteResponseStream)
if !ok {