diff --git a/app/rest/admin.go b/app/rest/admin.go index 7c2c0a2e..d768ca35 100644 --- a/app/rest/admin.go +++ b/app/rest/admin.go @@ -11,6 +11,7 @@ import ( "github.com/umputun/remark/app/store" ) +// admin provides router for all requests available for admin only type admin struct { dataStore store.Interface exporter migrator.Exporter diff --git a/app/rest/auth/auth.go b/app/rest/auth/auth.go index 1facf246..45b7e342 100644 --- a/app/rest/auth/auth.go +++ b/app/rest/auth/auth.go @@ -64,7 +64,7 @@ func (p Provider) LoginHandler(w http.ResponseWriter, r *http.Request) { state := randToken() session, err := p.Get(r, "remark") if err != nil { - log.Printf("[WARN] %s", err) + log.Printf("[DEBUG] can't get session, %s", err) } session.Values["state"] = state @@ -80,12 +80,14 @@ func (p Provider) LoginHandler(w http.ResponseWriter, r *http.Request) { } // return login url - log.Printf("[DEBUG] login url %s", p.conf.AuthCodeURL(state)) - http.Redirect(w, r, p.conf.AuthCodeURL(state), http.StatusTemporaryRedirect) + loginURL := p.conf.AuthCodeURL(state) + log.Printf("[DEBUG] login url %s", loginURL) + http.Redirect(w, r, loginURL, http.StatusTemporaryRedirect) } -// AuthHandler fills user info and redirects to "from" url +// AuthHandler fills user info and redirects to "from" url. This is callback url redirected locally by browser func (p Provider) AuthHandler(w http.ResponseWriter, r *http.Request) { + session, err := p.Get(r, "remark") if err != nil { http.Error(w, fmt.Sprintf("failed to get session, %s", err), http.StatusInternalServerError) @@ -140,7 +142,7 @@ func (p Provider) AuthHandler(w http.ResponseWriter, r *http.Request) { log.Printf("[DEBUG] %+v", jData) - // redirect to back url if presented + // redirect to back url if presented in login query params if fromURL, ok := session.Values["from"]; ok { http.Redirect(w, r, fromURL.(string), http.StatusTemporaryRedirect) return @@ -157,7 +159,10 @@ func (p Provider) LogoutHandler(w http.ResponseWriter, r *http.Request) { return } + session.Values["uinfo"] = "" session.Values["from"] = "" + session.Values["state"] = "" + delete(session.Values, "uinfo") delete(session.Values, "from") delete(session.Values, "state") diff --git a/app/rest/middleware.go b/app/rest/middleware.go index 1e16011b..577f499e 100644 --- a/app/rest/middleware.go +++ b/app/rest/middleware.go @@ -1,15 +1,21 @@ package rest import ( + "bytes" "context" + "fmt" + "io/ioutil" "log" "net/http" + "net/url" "os" + "regexp" "runtime/debug" "strings" "time" "github.com/didip/tollbooth" + "github.com/go-chi/chi/middleware" "github.com/go-chi/render" "github.com/go-errors/errors" "github.com/gorilla/sessions" @@ -110,13 +116,29 @@ func Recoverer(next http.Handler) http.Handler { type contextKey string +const ( + anonymous = iota + developer + full +) + // Auth adds auth from session and populate user info -func Auth(sessionStore *sessions.FilesystemStore, devMode bool, admins []string) func(http.Handler) http.Handler { +func Auth(sessionStore *sessions.FilesystemStore, admins []string, modes ...int) func(http.Handler) http.Handler { + + inModes := func(mode int) bool { + for _, m := range modes { + if m == mode { + return true + } + } + return false + } + f := func(h http.Handler) http.Handler { fn := func(w http.ResponseWriter, r *http.Request) { // for dev mode skip all real auth, make dev admin user - if devMode { + if inModes(developer) { user := store.User{ ID: "dev", Name: "developer one", @@ -138,22 +160,24 @@ func Auth(sessionStore *sessions.FilesystemStore, devMode bool, admins []string) } uinfoData, ok := session.Values["uinfo"] - if !ok { + if !ok && inModes(full) { http.Error(w, "Unauthorized", http.StatusUnauthorized) return } - user := uinfoData.(store.User) - for _, admin := range admins { - if admin == user.ID { - user.Admin = true - break + + if ok { + user := uinfoData.(store.User) + for _, admin := range admins { + if admin == user.ID { + user.Admin = true + break + } } + + ctx := r.Context() + ctx = context.WithValue(ctx, contextKey("user"), user) + r = r.WithContext(ctx) } - - ctx := r.Context() - ctx = context.WithValue(ctx, contextKey("user"), user) - r = r.WithContext(ctx) - h.ServeHTTP(w, r) } return http.HandlerFunc(fn) @@ -195,3 +219,88 @@ func GetUserInfo(r *http.Request) (user store.User, err error) { return store.User{}, errors.New("user can't be parsed") } + +// LoggerFlag type +type LoggerFlag int + +// logger flags enum +const ( + LogAll LoggerFlag = iota + LogUser + LogBody +) +const maxBody = 1024 + +var reMultWhtsp = regexp.MustCompile(`[\s\p{Zs}]{2,}`) + +// Logger middleware prints http log. Customized by set of LoggerFlag +func Logger(flags ...LoggerFlag) func(http.Handler) http.Handler { + + inFlags := func(f LoggerFlag) bool { + for _, flg := range flags { + if flg == LogAll || flg == f { + return true + } + } + return false + } + + f := func(h http.Handler) http.Handler { + + fn := func(w http.ResponseWriter, r *http.Request) { + ww := middleware.NewWrapResponseWriter(w, 1) + + body, user := func() (body string, user string) { + ctx := r.Context() + if ctx == nil { + return "", "" + } + + if inFlags(LogBody) { + if content, err := ioutil.ReadAll(r.Body); err == nil { + body = string(content) + r.Body = ioutil.NopCloser(bytes.NewReader(content)) + + if len(body) > 0 { + body = strings.Replace(body, "\n", " ", -1) + body = reMultWhtsp.ReplaceAllString(body, " ") + } + + if len(body) > maxBody { + body = body[:maxBody] + "..." + } + } + } + + if inFlags(LogUser) { + u, err := GetUserInfo(r) + if err == nil && u.Name != "" { + user = fmt.Sprintf(" - %s %q", u.ID, u.Name) + } + } + + return body, user + }() + + t1 := time.Now() + defer func() { + t2 := time.Now() + + q := r.URL.String() + if qun, err := url.QueryUnescape(q); err == nil { + q = qun + } + + log.Printf("[INFO] REST %s%s - %s - %s - %d (%d) - %v %s", + r.Method, user, q, strings.Split(r.RemoteAddr, ":")[0], + ww.Status(), ww.BytesWritten(), t2.Sub(t1), body) + }() + + h.ServeHTTP(ww, r) + } + return http.HandlerFunc(fn) + } + + return f + +} diff --git a/app/rest/server.go b/app/rest/server.go index 444338e3..1503c88d 100644 --- a/app/rest/server.go +++ b/app/rest/server.go @@ -39,10 +39,19 @@ type Server struct { func (s *Server) Run() { log.Print("[INFO] activate rest server") + applyDevMode := func(mode int) (modes []int) { + modes = append(modes, mode) + if s.DevMode { + modes = append(modes, developer) + } + return modes + } + router := chi.NewRouter() router.Use(middleware.RealIP, Recoverer) router.Use(middleware.Throttle(1000), middleware.Timeout(60*time.Second)) - router.Use(Limiter(10), AppInfo("remark", s.Version), Ping) + router.Use(Auth(s.SessionStore, s.Admins, applyDevMode(anonymous)...)) + router.Use(Limiter(10), AppInfo("remark", s.Version), Ping, Logger(LogAll)) router.Get("/login/google", s.AuthGoogle.LoginHandler) router.Get("/auth/google", s.AuthGoogle.AuthHandler) @@ -56,7 +65,7 @@ func (s *Server) Run() { rapi.Get("/last/{max}", s.lastCommentsCtrl) rapi.Get("/count", s.countCtrl) - rapi.With(Auth(s.SessionStore, s.DevMode, s.Admins)).Group(func(rauth chi.Router) { + rapi.With(Auth(s.SessionStore, s.Admins, applyDevMode(full)...)).Group(func(rauth chi.Router) { rauth.Post("/comment", s.createCommentCtrl) rauth.Get("/user", s.userInfoCtrl) rauth.Put("/vote/{id}", s.voteCtrl) @@ -112,7 +121,7 @@ func (s *Server) createCommentCtrl(w http.ResponseWriter, r *http.Request) { comment.User = user comment.User.IP = strings.Split(r.RemoteAddr, ":")[0] - log.Printf("[INFO] create comment %+v", comment) + log.Printf("[DEBUG] create comment %+v", comment) // check if user blocked if s.mod.checkBlocked(store.Locator{}, comment.User) { @@ -136,7 +145,7 @@ func (s *Server) createCommentCtrl(w http.ResponseWriter, r *http.Request) { func (s *Server) deleteCommentCtrl(w http.ResponseWriter, r *http.Request) { id := chi.URLParam(r, "id") - log.Printf("[INFO] delete comment %s", id) + log.Printf("[DEBUG] delete comment %s", id) url := r.URL.Query().Get("url") err := s.Store.Delete(store.Locator{URL: url}, id) @@ -153,7 +162,7 @@ func (s *Server) deleteCommentCtrl(w http.ResponseWriter, r *http.Request) { // GET /find?url=post-url func (s *Server) findCommentsCtrl(w http.ResponseWriter, r *http.Request) { url := r.URL.Query().Get("url") - log.Printf("[INFO] get comments for %s", url) + log.Printf("[DEBUG] get comments for %s", url) comments, err := s.Store.Find(store.Request{Locator: store.Locator{URL: url}}) if err != nil { @@ -190,7 +199,7 @@ func (s *Server) commentByIDCtrl(w http.ResponseWriter, r *http.Request) { id := chi.URLParam(r, "id") url := r.URL.Query().Get("url") - log.Printf("[INFO] get comments by id %s, %s", id, url) + log.Printf("[DEBUG] get comments by id %s, %s", id, url) comment, err := s.Store.Get(store.Locator{URL: url}, id) if err != nil { @@ -233,7 +242,7 @@ func (s *Server) voteCtrl(w http.ResponseWriter, r *http.Request) { } id := chi.URLParam(r, "id") - log.Printf("[INFO] vote for comment %s", id) + log.Printf("[DEBUG] vote for comment %s", id) url := r.URL.Query().Get("url") vote := r.URL.Query().Get("vote") == "1" diff --git a/app/store/bolt.go b/app/store/bolt.go index e21c8712..8dda970d 100644 --- a/app/store/bolt.go +++ b/app/store/bolt.go @@ -13,9 +13,8 @@ import ( ) // BoltDB implements store.Interface. Each instance represents one site. -// Keys are commendID. Each url (post) makes it's own bucket. -// In addition there is a bucket "last" with reference to other buckets+keys to all cross-posts last comment extraction. -// Thread safe. +// Keys are commentID. Each url (post) makes it's own bucket. In addition there is a bucket "last" with +// reference to other buckets+keys to all cross-posts last comment extraction. Thread safe. type BoltDB struct { *bolt.DB } @@ -54,7 +53,7 @@ func (b *BoltDB) Create(comment Comment) (string, error) { err := b.Update(func(tx *bolt.Tx) error { bucket, e := tx.CreateBucketIfNotExists([]byte(comment.Locator.URL)) if e != nil { - return errors.Wrapf(e, "can't make bucket", comment.Locator.URL) + return errors.Wrapf(e, "can't make or open bucket", comment.Locator.URL) } // check if key already in store, reject doubles @@ -250,6 +249,7 @@ func (b *BoltDB) Vote(locator Locator, commentID string, userID string, val bool // update votes and score comment.Votes[userID] = val + if val { comment.Score++ } else { @@ -321,10 +321,6 @@ func (b *BoltDB) IsBlocked(locator Locator, userID string) (result bool) { return result } -func (b *BoltDB) bucketForBlock(locator Locator, userID string) []byte { - return []byte(fmt.Sprintf("%s%s", blocksBucketPrefix, locator.SiteID)) -} - // List returns list of buckets, which is list of all commented posts func (b BoltDB) List(locator Locator) (result []string, err error) { @@ -339,11 +335,17 @@ func (b BoltDB) List(locator Locator) (result []string, err error) { return result, err } +func (b *BoltDB) bucketForBlock(locator Locator, userID string) []byte { + return []byte(fmt.Sprintf("%s%s", blocksBucketPrefix, locator.SiteID)) +} + +// ref represents key:value pair for extra, index-only buckets type ref struct { key string value string } +// refFromComment makes reference record used for related buckets referencing prim data set func refFromComment(comment Comment) *ref { result := ref{ key: fmt.Sprintf("%s!!%s", comment.Timestamp.Format(time.RFC3339Nano), comment.ID), diff --git a/app/store/store_test.go b/app/store/store_test.go new file mode 100644 index 00000000..f9d05baf --- /dev/null +++ b/app/store/store_test.go @@ -0,0 +1,39 @@ +package store + +import "testing" +import "github.com/stretchr/testify/assert" + +func TestStore_MakeCommentID(t *testing.T) { + cid1 := makeCommentID() + assert.True(t, len(cid1) > 8, "cid1 is long enough") + + cid2 := makeCommentID() + assert.True(t, len(cid2) > 8, "cid2 is long enough") + + assert.NotEqual(t, cid1, cid2, "cids different") +} + +func TestStore_SanitizeComment(t *testing.T) { + + tbl := []struct { + inp Comment + out Comment + }{ + {inp: Comment{}, out: Comment{}}, + { + inp: Comment{ + Text: `blah XSS` + "\n\t", + User: User{ID: `username`}, + }, + out: Comment{ + Text: `blah XSS`, + User: User{ID: `<a href="http://blah.com">username</a>`}, + }, + }, + } + + for n, tt := range tbl { + out := sanitizeComment(tt.inp) + assert.Equal(t, tt.out, out, "check #%d", n) + } +}