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)
+ }
+}