diff options
| -rw-r--r-- | internal/api/api.go | 9 | ||||
| -rw-r--r-- | internal/api/errorcode_string.go | 9 | ||||
| -rw-r--r-- | internal/cryptoutil/cryptoutil.go | 26 | ||||
| -rw-r--r-- | internal/cryptoutil/cryptoutil_test.go | 20 | ||||
| -rw-r--r-- | roseh.moe.go | 138 | ||||
| -rw-r--r-- | tools/notes/notes.go | 108 |
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 { |
