Compare commits
14 Commits
91f5a63366
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5aea734d2b | ||
|
|
69658a7bd4 | ||
|
|
577f902ab3 | ||
|
|
ec5ef9a5c1 | ||
|
|
79d465bd7b | ||
|
|
c3294de4c6 | ||
|
|
fc07515dbd | ||
|
|
8d3c025660 | ||
|
|
126ccfa715 | ||
|
|
b6cdb2b860 | ||
|
|
94345eff97 | ||
|
|
894d152aae | ||
|
|
dd4c74598d | ||
|
|
2824e51c29 |
5
cmd/web/context.go
Normal file
5
cmd/web/context.go
Normal file
@@ -0,0 +1,5 @@
|
||||
package main
|
||||
|
||||
type contextKey string
|
||||
|
||||
const isAuthenticatedContextKey = contextKey("isAuthenticated")
|
||||
@@ -4,11 +4,14 @@ import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"gitea.local.lab/Lbenedar/snippetbox/internal/models"
|
||||
"gitea.local.lab/Lbenedar/snippetbox/internal/validator"
|
||||
"gitea.local.lab/Lbenedar/snippetbox/ui"
|
||||
"github.com/julienschmidt/httprouter"
|
||||
)
|
||||
|
||||
@@ -19,6 +22,33 @@ type snippetCreateForm struct {
|
||||
validator.Validator `form:"-"`
|
||||
}
|
||||
|
||||
type userSignupForm struct {
|
||||
Name string `form:"name"`
|
||||
Email string `form:"email"`
|
||||
Password string `form:"password"`
|
||||
validator.Validator `form:"-"`
|
||||
}
|
||||
|
||||
type userLoginForm struct {
|
||||
Email string `form:"email"`
|
||||
Password string `form:"password"`
|
||||
validator.Validator `form:"-"`
|
||||
}
|
||||
|
||||
type accountViewForm struct {
|
||||
Name string `form:"name"`
|
||||
Email string `form:"email"`
|
||||
Joined string `form:"joined"`
|
||||
validator.Validator `form:"-"`
|
||||
}
|
||||
|
||||
type passwordChangeForm struct {
|
||||
CurrentPassword string `form:"curr_pass"`
|
||||
NewPassword string `form:"new_pass"`
|
||||
ConfirmPassword string `form:"conf_pass"`
|
||||
validator.Validator `form:"-"`
|
||||
}
|
||||
|
||||
func (app *application) render(w http.ResponseWriter, status int, page string, data *templateData) {
|
||||
ts, ok := app.templateCache[page]
|
||||
if !ok {
|
||||
@@ -98,7 +128,7 @@ func (app *application) snippetCreatePost(w http.ResponseWriter, r *http.Request
|
||||
form.CheckField(validator.NotBlank(form.Title), "title", "This field cannot be blank")
|
||||
form.CheckField(validator.MaxChars(form.Title, 100), "title", "This field cannot be more than 100 character long")
|
||||
form.CheckField(validator.NotBlank(form.Content), "content", "This field cannot be blank")
|
||||
form.CheckField(validator.PermittedInt(form.Expires, 1, 7, 365), "expires", "This field must equal 1, 7 and 365")
|
||||
form.CheckField(validator.PermittedValue(form.Expires, 1, 7, 365), "expires", "This field must equal 1, 7 and 365")
|
||||
|
||||
if !form.Valid() {
|
||||
data := app.newTemplateData(r)
|
||||
@@ -119,21 +149,186 @@ func (app *application) snippetCreatePost(w http.ResponseWriter, r *http.Request
|
||||
}
|
||||
|
||||
func (app *application) userSignupPost(w http.ResponseWriter, r *http.Request) {
|
||||
fmt.Fprintln(w, "Create a new user...")
|
||||
var form userSignupForm
|
||||
err := app.decodePostForm(r, &form)
|
||||
|
||||
if err != nil {
|
||||
app.notFound(w)
|
||||
return
|
||||
}
|
||||
|
||||
form.CheckField(validator.NotBlank(form.Name), "name", "This field cannot be blank")
|
||||
form.CheckField(validator.NotBlank(form.Email), "email", "This field cannot be blank")
|
||||
form.CheckField(validator.Matches(form.Email, validator.EmailRX), "email", "This field must be a valid email address")
|
||||
form.CheckField(validator.NotBlank(form.Password), "password", "This field cannot be blank")
|
||||
form.CheckField(validator.MinChars(form.Password, 8), "password", "This field must be at least 8 character long")
|
||||
|
||||
if !form.Valid() {
|
||||
data := app.newTemplateData(r)
|
||||
data.Form = form
|
||||
app.render(w, http.StatusUnprocessableEntity, "signup.tmpl", data)
|
||||
return
|
||||
}
|
||||
|
||||
err = app.users.Insert(form.Name, form.Email, form.Password)
|
||||
if err != nil {
|
||||
if errors.Is(err, models.ErrDuplicateEmail) {
|
||||
form.AddFieldError("email", "Email address already in use")
|
||||
|
||||
data := app.newTemplateData(r)
|
||||
data.Form = form
|
||||
app.render(w, http.StatusUnprocessableEntity, "signup.tmpl", data)
|
||||
} else {
|
||||
app.serverError(w, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
app.sessionManager.Put(r.Context(), "flash", "Your singup was successful. Please log in.")
|
||||
http.Redirect(w, r, "/user/login", http.StatusSeeOther)
|
||||
}
|
||||
|
||||
func (app *application) userLoginPost(w http.ResponseWriter, r *http.Request) {
|
||||
fmt.Fprintln(w, "Login the user...")
|
||||
var form userLoginForm
|
||||
|
||||
err := app.decodePostForm(r, &form)
|
||||
if err != nil {
|
||||
app.clientError(w, http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
form.CheckField(validator.NotBlank(form.Email), "email", "This field cannot be blank")
|
||||
form.CheckField(validator.Matches(form.Email, validator.EmailRX), "email", "This field must be a valid email address")
|
||||
form.CheckField(validator.NotBlank(form.Password), "password", "This field cannot be blank")
|
||||
|
||||
if !form.Valid() {
|
||||
data := app.newTemplateData(r)
|
||||
data.Form = form
|
||||
app.render(w, http.StatusUnprocessableEntity, "login.tmpl", data)
|
||||
return
|
||||
}
|
||||
|
||||
id, err := app.users.Authenticate(form.Email, form.Password)
|
||||
if err != nil {
|
||||
if errors.Is(err, models.ErrInvalidCredentials) {
|
||||
form.AddNonFieldError("Email or password is incorrect")
|
||||
|
||||
data := app.newTemplateData(r)
|
||||
data.Form = form
|
||||
app.render(w, http.StatusUnprocessableEntity, "login.tmpl", data)
|
||||
} else {
|
||||
app.serverError(w, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
err = app.sessionManager.RenewToken(r.Context())
|
||||
if err != nil {
|
||||
app.serverError(w, err)
|
||||
return
|
||||
}
|
||||
|
||||
app.sessionManager.Put(r.Context(), "authenticatedUserID", id)
|
||||
redirectPath := app.sessionManager.PopString(r.Context(), "redirectParent")
|
||||
if redirectPath != "" {
|
||||
http.Redirect(w, r, redirectPath, http.StatusSeeOther)
|
||||
return
|
||||
}
|
||||
http.Redirect(w, r, "/snippet/create", http.StatusSeeOther)
|
||||
}
|
||||
|
||||
func (app *application) userLogoutPost(w http.ResponseWriter, r *http.Request) {
|
||||
fmt.Fprintln(w, "Logout the user...")
|
||||
err := app.sessionManager.RenewToken(r.Context())
|
||||
if err != nil {
|
||||
app.serverError(w, err)
|
||||
return
|
||||
}
|
||||
app.sessionManager.Remove(r.Context(), "authenticatedUserID")
|
||||
app.sessionManager.Put(r.Context(), "flash", "You've been logged out succefully!")
|
||||
|
||||
http.Redirect(w, r, "/", http.StatusSeeOther)
|
||||
}
|
||||
|
||||
func (app *application) userSignup(w http.ResponseWriter, r *http.Request) {
|
||||
fmt.Fprintln(w, "Display a HTML for signing up a new user...")
|
||||
data := app.newTemplateData(r)
|
||||
data.Form = userSignupForm{}
|
||||
app.render(w, http.StatusOK, "signup.tmpl", data)
|
||||
}
|
||||
|
||||
func (app *application) userLogin(w http.ResponseWriter, r *http.Request) {
|
||||
fmt.Fprintln(w, "Display a HTML for log in a user...")
|
||||
data := app.newTemplateData(r)
|
||||
data.Form = userLoginForm{}
|
||||
app.render(w, http.StatusOK, "login.tmpl", data)
|
||||
}
|
||||
|
||||
func ping(w http.ResponseWriter, r *http.Request) {
|
||||
w.Write([]byte("OK"))
|
||||
}
|
||||
|
||||
func (app *application) about(w http.ResponseWriter, r *http.Request) {
|
||||
data := app.newTemplateData(r)
|
||||
text, err := fs.ReadFile(ui.Files, "static/text/about.txt")
|
||||
if err != nil {
|
||||
app.serverError(w, err)
|
||||
return
|
||||
}
|
||||
data.AboutText = string(text)
|
||||
app.render(w, http.StatusOK, "about.tmpl", data)
|
||||
}
|
||||
|
||||
func (app *application) accountView(w http.ResponseWriter, r *http.Request) {
|
||||
data := app.newTemplateData(r)
|
||||
id := app.sessionManager.Get(r.Context(), "authenticatedUserID").(int)
|
||||
|
||||
user, err := app.users.GetById(id)
|
||||
if err != nil {
|
||||
app.serverError(w, err)
|
||||
return
|
||||
}
|
||||
data.Form = accountViewForm{Name: user.Name, Email: user.Email, Joined: humanDate(user.Created)}
|
||||
app.render(w, http.StatusOK, "account.tmpl", data)
|
||||
}
|
||||
|
||||
func (app *application) accountChangePassword(w http.ResponseWriter, r *http.Request) {
|
||||
data := app.newTemplateData(r)
|
||||
data.Form = passwordChangeForm{}
|
||||
app.render(w, http.StatusOK, "password.tmpl", data)
|
||||
}
|
||||
|
||||
func (app *application) accountChangePasswordPost(w http.ResponseWriter, r *http.Request) {
|
||||
form := passwordChangeForm{}
|
||||
err := app.decodePostForm(r, &form)
|
||||
if err != nil {
|
||||
app.clientError(w, http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
form.CheckField(validator.NotBlank(form.CurrentPassword), "curr_pass", "This field cannot be blank")
|
||||
form.CheckField(validator.NotBlank(form.NewPassword), "new_pass", "This field cannot be blank")
|
||||
form.CheckField(validator.MinChars(form.NewPassword, 8), "new_pass", "Minimum password length is 8 symbols")
|
||||
form.CheckField(validator.NotBlank(form.ConfirmPassword), "conf_pass", "This field cannot be blank")
|
||||
form.CheckField(strings.Compare(form.NewPassword, form.ConfirmPassword) == 0, "conf_pass", "Confirm password does not match")
|
||||
|
||||
if !form.Valid() {
|
||||
data := app.newTemplateData(r)
|
||||
data.Form = form
|
||||
app.render(w, http.StatusUnprocessableEntity, "password.tmpl", data)
|
||||
return
|
||||
}
|
||||
|
||||
id := app.sessionManager.Get(r.Context(), "authenticatedUserID").(int)
|
||||
err = app.users.ChangePassword(id, form.CurrentPassword, form.NewPassword)
|
||||
if err != nil {
|
||||
if errors.Is(err, models.ErrInvalidCredentials) {
|
||||
form.AddNonFieldError("Password is incorrect")
|
||||
|
||||
data := app.newTemplateData(r)
|
||||
data.Form = form
|
||||
app.render(w, http.StatusUnprocessableEntity, "password.tmpl", data)
|
||||
} else {
|
||||
app.serverError(w, err)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
app.sessionManager.Put(r.Context(), "flash", "Password successfully changed!")
|
||||
http.Redirect(w, r, "/account/view", http.StatusSeeOther)
|
||||
}
|
||||
|
||||
228
cmd/web/handlers_test.go
Normal file
228
cmd/web/handlers_test.go
Normal file
@@ -0,0 +1,228 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/url"
|
||||
"testing"
|
||||
|
||||
"gitea.local.lab/Lbenedar/snippetbox/internal/assert"
|
||||
)
|
||||
|
||||
func TestPing(t *testing.T) {
|
||||
app := newTestApplication(t)
|
||||
|
||||
ts := newTestServer(t, app.routes())
|
||||
defer ts.Close()
|
||||
|
||||
statusCode, _, body := ts.get(t, "/ping")
|
||||
|
||||
assert.Equal(t, statusCode, http.StatusOK)
|
||||
assert.Equal(t, body, "OK")
|
||||
}
|
||||
|
||||
func TestSnippetView(t *testing.T) {
|
||||
app := newTestApplication(t)
|
||||
|
||||
ts := newTestServer(t, app.routes())
|
||||
defer ts.Close()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
urlPath string
|
||||
wantCode int
|
||||
wantBody string
|
||||
}{
|
||||
{
|
||||
name: "Valid ID",
|
||||
urlPath: "/snippet/view/1",
|
||||
wantCode: http.StatusOK,
|
||||
wantBody: "An old silent pond...",
|
||||
},
|
||||
{
|
||||
name: "Non-existend ID",
|
||||
urlPath: "/snippet/view/2",
|
||||
wantCode: http.StatusNotFound,
|
||||
},
|
||||
{
|
||||
name: "Negative ID",
|
||||
urlPath: "/snippet/view/-1",
|
||||
wantCode: http.StatusNotFound,
|
||||
},
|
||||
{
|
||||
name: "Decimal ID",
|
||||
urlPath: "/snippet/view/1.23",
|
||||
wantCode: http.StatusNotFound,
|
||||
},
|
||||
{
|
||||
name: "String ID",
|
||||
urlPath: "/snippet/view/foo",
|
||||
wantCode: http.StatusNotFound,
|
||||
},
|
||||
{
|
||||
name: "Empty ID",
|
||||
urlPath: "/snippet/view/",
|
||||
wantCode: http.StatusNotFound,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
code, _, body := ts.get(t, tt.urlPath)
|
||||
|
||||
assert.Equal(t, code, tt.wantCode)
|
||||
|
||||
if tt.wantBody != "" {
|
||||
assert.StringContains(t, body, tt.wantBody)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUserSignup(t *testing.T) {
|
||||
app := newTestApplication(t)
|
||||
ts := newTestServer(t, app.routes())
|
||||
defer ts.Close()
|
||||
|
||||
_, _, body := ts.get(t, "/user/signup")
|
||||
validCsrfToken := extractCSRFToken(t, body)
|
||||
t.Logf("CsrfToken: %s", validCsrfToken)
|
||||
|
||||
const (
|
||||
validName = "Bob"
|
||||
validPassword = "validPa$$word"
|
||||
validEmail = "bob@example.com"
|
||||
formTag = "<form action='/user/signup' method='POST' novalidate>"
|
||||
)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
userName string
|
||||
userEmail string
|
||||
userPassword string
|
||||
csrfToken string
|
||||
wantCode int
|
||||
wantFormTag string
|
||||
}{
|
||||
{
|
||||
name: "Valid submission",
|
||||
userName: validName,
|
||||
userEmail: validEmail,
|
||||
userPassword: validPassword,
|
||||
csrfToken: validCsrfToken,
|
||||
wantCode: http.StatusSeeOther,
|
||||
},
|
||||
{
|
||||
name: "Invalid CSRF Token",
|
||||
userName: validName,
|
||||
userEmail: validEmail,
|
||||
userPassword: validPassword,
|
||||
csrfToken: "wrongToken",
|
||||
wantCode: http.StatusBadRequest,
|
||||
},
|
||||
{
|
||||
name: "Empty name",
|
||||
userName: "",
|
||||
userEmail: validEmail,
|
||||
userPassword: validPassword,
|
||||
csrfToken: validCsrfToken,
|
||||
wantCode: http.StatusUnprocessableEntity,
|
||||
wantFormTag: formTag,
|
||||
},
|
||||
{
|
||||
name: "Empty email",
|
||||
userName: validName,
|
||||
userEmail: "",
|
||||
userPassword: validPassword,
|
||||
csrfToken: validCsrfToken,
|
||||
wantCode: http.StatusUnprocessableEntity,
|
||||
wantFormTag: formTag,
|
||||
},
|
||||
{
|
||||
name: "Empty password",
|
||||
userName: validName,
|
||||
userEmail: validEmail,
|
||||
userPassword: "",
|
||||
csrfToken: validCsrfToken,
|
||||
wantCode: http.StatusUnprocessableEntity,
|
||||
wantFormTag: formTag,
|
||||
},
|
||||
{
|
||||
name: "Invalid email",
|
||||
userName: validName,
|
||||
userEmail: "bob@example.",
|
||||
userPassword: validPassword,
|
||||
csrfToken: validCsrfToken,
|
||||
wantCode: http.StatusUnprocessableEntity,
|
||||
wantFormTag: formTag,
|
||||
},
|
||||
{
|
||||
name: "Short password",
|
||||
userName: validName,
|
||||
userEmail: validEmail,
|
||||
userPassword: "pa$$",
|
||||
csrfToken: validCsrfToken,
|
||||
wantCode: http.StatusUnprocessableEntity,
|
||||
wantFormTag: formTag,
|
||||
},
|
||||
{
|
||||
name: "Duplicate email",
|
||||
userName: validName,
|
||||
userEmail: "dupe@example.com",
|
||||
userPassword: validPassword,
|
||||
csrfToken: validCsrfToken,
|
||||
wantCode: http.StatusUnprocessableEntity,
|
||||
wantFormTag: formTag,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
form := url.Values{}
|
||||
form.Add("name", tt.userName)
|
||||
form.Add("email", tt.userEmail)
|
||||
form.Add("password", tt.userPassword)
|
||||
form.Add("csrf_token", tt.csrfToken)
|
||||
|
||||
code, _, body := ts.postForm(t, "/user/signup", form)
|
||||
|
||||
assert.Equal(t, code, tt.wantCode)
|
||||
if tt.wantFormTag != "" {
|
||||
assert.StringContains(t, body, tt.wantFormTag)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSnippetCreate(t *testing.T) {
|
||||
app := newTestApplication(t)
|
||||
ts := newTestServer(t, app.routes())
|
||||
defer ts.Close()
|
||||
|
||||
const (
|
||||
validPassword = "validPa$$word"
|
||||
validEmail = "bob@example.com"
|
||||
)
|
||||
|
||||
t.Run("Unauthenticated", func(t *testing.T) {
|
||||
code, head, _ := ts.get(t, "/snippet/create")
|
||||
assert.Equal(t, code, http.StatusSeeOther)
|
||||
assert.Equal(t, head.Get("Location"), "/user/login")
|
||||
})
|
||||
|
||||
t.Run("Authenticated", func(t *testing.T) {
|
||||
_, _, body := ts.get(t, "/user/login")
|
||||
validCsrfToken := extractCSRFToken(t, body)
|
||||
t.Logf("CsrfToken: %s", validCsrfToken)
|
||||
|
||||
form := url.Values{}
|
||||
form.Add("email", validEmail)
|
||||
form.Add("password", validPassword)
|
||||
form.Add("csrf_token", validCsrfToken)
|
||||
code, _, _ := ts.postForm(t, "/user/login", form)
|
||||
t.Logf("Auth result: %d", code)
|
||||
|
||||
code, _, body = ts.get(t, "/snippet/create")
|
||||
assert.Equal(t, code, http.StatusOK)
|
||||
assert.StringContains(t, body, "<form action='/snippet/create' method='POST'>")
|
||||
})
|
||||
}
|
||||
@@ -13,7 +13,13 @@ func (app *application) serverError(w http.ResponseWriter, err error) {
|
||||
trace := fmt.Sprintf("%s\n%s", err.Error(), debug.Stack())
|
||||
app.errorLog.Output(2, trace)
|
||||
|
||||
http.Error(w, http.StatusText(http.StatusInternalServerError), http.StatusInternalServerError)
|
||||
var errorText string
|
||||
if app.debug {
|
||||
errorText = trace
|
||||
} else {
|
||||
errorText = http.StatusText(http.StatusInternalServerError)
|
||||
}
|
||||
http.Error(w, errorText, http.StatusInternalServerError)
|
||||
}
|
||||
|
||||
func (app *application) clientError(w http.ResponseWriter, status int) {
|
||||
@@ -24,13 +30,13 @@ func (app *application) notFound(w http.ResponseWriter) {
|
||||
app.clientError(w, http.StatusNotFound)
|
||||
}
|
||||
|
||||
func (app *application) decodePostForm(r *http.Request, snippetForm *snippetCreateForm) error {
|
||||
func (app *application) decodePostForm(r *http.Request, templateForm any) error {
|
||||
err := r.ParseForm()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
err = app.formDecoder.Decode(&snippetForm, r.PostForm)
|
||||
err = app.formDecoder.Decode(&templateForm, r.PostForm)
|
||||
if err != nil {
|
||||
var invalidDecoderError *form.InvalidDecoderError
|
||||
if errors.As(err, &invalidDecoderError) {
|
||||
@@ -40,3 +46,11 @@ func (app *application) decodePostForm(r *http.Request, snippetForm *snippetCrea
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (app *application) isAuthenticated(r *http.Request) bool {
|
||||
isAuthenticated, ok := r.Context().Value(isAuthenticatedContextKey).(bool)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return isAuthenticated
|
||||
}
|
||||
|
||||
@@ -21,17 +21,19 @@ import (
|
||||
type application struct {
|
||||
errorLog *log.Logger
|
||||
infoLog *log.Logger
|
||||
snippets *models.SnippetModel
|
||||
users *models.UserModel
|
||||
snippets models.SnippetModelInterface
|
||||
users models.UserModelInterface
|
||||
templateCache map[string]*template.Template
|
||||
formDecoder *form.Decoder
|
||||
sessionManager *scs.SessionManager
|
||||
debug bool
|
||||
}
|
||||
|
||||
type config struct {
|
||||
addr string
|
||||
staticDir string
|
||||
dsn string
|
||||
debug bool
|
||||
}
|
||||
|
||||
func main() {
|
||||
@@ -39,6 +41,7 @@ func main() {
|
||||
flag.StringVar(&cfg.addr, "addr", ":4000", "HTTP network address")
|
||||
flag.StringVar(&cfg.staticDir, "static-dir", "./ui/static", "Path to static assets")
|
||||
flag.StringVar(&cfg.dsn, "dsn", "web:pass@/snippetbox?parseTime=true", "MySQL data source name")
|
||||
flag.BoolVar(&cfg.debug, "debug", false, "Option to enable debug mode")
|
||||
flag.Parse()
|
||||
|
||||
errorLog := log.New(os.Stderr, "ERROR\t", log.Ldate|log.Ltime|log.Lshortfile)
|
||||
@@ -68,6 +71,7 @@ func main() {
|
||||
templateCache: templateCache,
|
||||
formDecoder: formDecoder,
|
||||
sessionManager: sessionManager,
|
||||
debug: cfg.debug,
|
||||
}
|
||||
|
||||
tlsConfig := &tls.Config{
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
|
||||
"github.com/justinas/nosurf"
|
||||
)
|
||||
|
||||
func secureHeaders(next http.Handler) http.Handler {
|
||||
@@ -10,14 +13,25 @@ func secureHeaders(next http.Handler) http.Handler {
|
||||
w.Header().Set("Content-Security-Policy",
|
||||
"default-src 'self'; style-src 'self' fonts.googleapis.com; font-src fonts.gstatic.com")
|
||||
w.Header().Set("Referrer-Policy", "origin-when-cross-origin")
|
||||
w.Header().Set("X-Content-Type-Options", "nosniff")
|
||||
w.Header().Set("X-Frame-Options", "deny")
|
||||
w.Header().Set("X-Content-Type-Options", "deny")
|
||||
w.Header().Set("X-Frame-Options", "nosniff")
|
||||
w.Header().Set("X-XSS-Protection", "0")
|
||||
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
func noSurf(next http.Handler) http.Handler {
|
||||
csrfHanlder := nosurf.New(next)
|
||||
csrfHanlder.SetBaseCookie(http.Cookie{
|
||||
HttpOnly: false,
|
||||
Path: "/",
|
||||
Secure: false,
|
||||
})
|
||||
|
||||
return csrfHanlder
|
||||
}
|
||||
|
||||
func (app *application) logRequest(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
app.infoLog.Printf("%s - %s %s %s", r.RemoteAddr, r.Proto, r.Method, r.URL.RequestURI())
|
||||
@@ -38,3 +52,37 @@ func (app *application) recoverPanic(next http.Handler) http.Handler {
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
func (app *application) requireAuthentication(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if !app.isAuthenticated(r) {
|
||||
app.sessionManager.Put(r.Context(), "redirectParent", r.URL.Path)
|
||||
http.Redirect(w, r, "/user/login", http.StatusSeeOther)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
func (app *application) authenticate(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
id := app.sessionManager.GetInt(r.Context(), "authenticatedUserID")
|
||||
if id == 0 {
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
exists, err := app.users.Exists(id)
|
||||
if err != nil {
|
||||
app.serverError(w, err)
|
||||
return
|
||||
}
|
||||
|
||||
if exists {
|
||||
ctx := context.WithValue(r.Context(), isAuthenticatedContextKey, true)
|
||||
r = r.WithContext(ctx)
|
||||
}
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
46
cmd/web/middleware_test.go
Normal file
46
cmd/web/middleware_test.go
Normal file
@@ -0,0 +1,46 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"gitea.local.lab/Lbenedar/snippetbox/internal/assert"
|
||||
)
|
||||
|
||||
func TestSecureHeaders(t *testing.T) {
|
||||
rr := httptest.NewRecorder()
|
||||
|
||||
r, err := http.NewRequest(http.MethodGet, "/", nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
next := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Write([]byte("OK"))
|
||||
})
|
||||
secureHeaders(next).ServeHTTP(rr, r)
|
||||
|
||||
rs := rr.Result()
|
||||
|
||||
expectedValue := "default-src 'self'; style-src 'self' fonts.googleapis.com; font-src fonts.gstatic.com"
|
||||
assert.Equal(t, rs.Header.Get("Content-Security-Policy"), expectedValue)
|
||||
|
||||
expectedValue = "nosniff"
|
||||
assert.Equal(t, rs.Header.Get("X-Frame-Options"), expectedValue)
|
||||
|
||||
expectedValue = "0"
|
||||
assert.Equal(t, rs.Header.Get("X-XSS-Protection"), expectedValue)
|
||||
|
||||
assert.Equal(t, rs.StatusCode, http.StatusOK)
|
||||
|
||||
defer rs.Body.Close()
|
||||
body, err := io.ReadAll(rs.Body)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
bytes.TrimSpace(body)
|
||||
|
||||
assert.Equal(t, string(body), "OK")
|
||||
}
|
||||
@@ -3,6 +3,8 @@ package main
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"gitea.local.lab/Lbenedar/snippetbox/ui"
|
||||
|
||||
"github.com/julienschmidt/httprouter"
|
||||
"github.com/justinas/alice"
|
||||
)
|
||||
@@ -14,21 +16,30 @@ func (app *application) routes() http.Handler {
|
||||
app.notFound(w)
|
||||
})
|
||||
|
||||
fileServer := http.FileServer(http.Dir("./ui/static/"))
|
||||
router.Handler(http.MethodGet, "/static/*filepath", http.StripPrefix("/static", fileServer))
|
||||
fileServer := http.FileServer(http.FS(ui.Files))
|
||||
router.Handler(http.MethodGet, "/static/*filepath", fileServer)
|
||||
|
||||
dynamic := alice.New(app.sessionManager.LoadAndSave)
|
||||
router.HandlerFunc(http.MethodGet, "/ping", ping)
|
||||
|
||||
dynamic := alice.New(app.sessionManager.LoadAndSave, noSurf, app.authenticate)
|
||||
|
||||
router.Handler(http.MethodGet, "/", dynamic.ThenFunc(app.home))
|
||||
router.Handler(http.MethodGet, "/about", dynamic.ThenFunc(app.about))
|
||||
router.Handler(http.MethodGet, "/snippet/view/:id", dynamic.ThenFunc(app.snippetView))
|
||||
router.Handler(http.MethodGet, "/snippet/create", dynamic.ThenFunc(app.snippetCreate))
|
||||
router.Handler(http.MethodPost, "/snippet/create", dynamic.ThenFunc(app.snippetCreatePost))
|
||||
|
||||
router.Handler(http.MethodGet, "/user/signup", dynamic.ThenFunc(app.userSignup))
|
||||
router.Handler(http.MethodPost, "/user/signup", dynamic.ThenFunc(app.userSignupPost))
|
||||
router.Handler(http.MethodGet, "/user/login", dynamic.ThenFunc(app.userLogin))
|
||||
router.Handler(http.MethodPost, "/user/login", dynamic.ThenFunc(app.userLoginPost))
|
||||
router.Handler(http.MethodPost, "/user/logout", dynamic.ThenFunc(app.userLogoutPost))
|
||||
|
||||
protected := dynamic.Append(app.requireAuthentication)
|
||||
|
||||
router.Handler(http.MethodGet, "/account/view", protected.ThenFunc(app.accountView))
|
||||
router.Handler(http.MethodGet, "/account/password", protected.ThenFunc(app.accountChangePassword))
|
||||
router.Handler(http.MethodPost, "/account/password", protected.ThenFunc(app.accountChangePasswordPost))
|
||||
router.Handler(http.MethodGet, "/snippet/create", protected.ThenFunc(app.snippetCreate))
|
||||
router.Handler(http.MethodPost, "/snippet/create", protected.ThenFunc(app.snippetCreatePost))
|
||||
router.Handler(http.MethodPost, "/user/logout", protected.ThenFunc(app.userLogoutPost))
|
||||
|
||||
standard := alice.New(app.recoverPanic, app.logRequest, secureHeaders)
|
||||
return standard.Then(router)
|
||||
|
||||
@@ -2,23 +2,32 @@ package main
|
||||
|
||||
import (
|
||||
"html/template"
|
||||
"io/fs"
|
||||
"net/http"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"gitea.local.lab/Lbenedar/snippetbox/internal/models"
|
||||
"gitea.local.lab/Lbenedar/snippetbox/ui"
|
||||
"github.com/justinas/nosurf"
|
||||
)
|
||||
|
||||
type templateData struct {
|
||||
CurrentYear int
|
||||
Snippet *models.Snippet
|
||||
Snippets []*models.Snippet
|
||||
Form any
|
||||
Flash string
|
||||
CurrentYear int
|
||||
Snippet *models.Snippet
|
||||
Snippets []*models.Snippet
|
||||
Form any
|
||||
Flash string
|
||||
IsAuthenticated bool
|
||||
CSRFToken string
|
||||
AboutText string
|
||||
}
|
||||
|
||||
func humanDate(t time.Time) string {
|
||||
return t.Format("02 Jan 2006 at 15:04")
|
||||
if t.IsZero() {
|
||||
return ""
|
||||
}
|
||||
return t.UTC().Format("02 Jan 2006 at 15:04")
|
||||
}
|
||||
|
||||
var functions = template.FuncMap{
|
||||
@@ -27,32 +36,30 @@ var functions = template.FuncMap{
|
||||
|
||||
func (app *application) newTemplateData(r *http.Request) *templateData {
|
||||
return &templateData{
|
||||
CurrentYear: time.Now().Year(),
|
||||
Flash: app.sessionManager.PopString(r.Context(), "flash"),
|
||||
CurrentYear: time.Now().Year(),
|
||||
Flash: app.sessionManager.PopString(r.Context(), "flash"),
|
||||
IsAuthenticated: app.isAuthenticated(r),
|
||||
CSRFToken: nosurf.Token(r),
|
||||
}
|
||||
}
|
||||
|
||||
func newTemplateCache() (map[string]*template.Template, error) {
|
||||
cache := map[string]*template.Template{}
|
||||
|
||||
pages, err := filepath.Glob("./ui/html/pages/*.tmpl")
|
||||
pages, err := fs.Glob(ui.Files, "html/pages/*.tmpl")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
for _, page := range pages {
|
||||
name := filepath.Base(page)
|
||||
ts, err := template.New(name).Funcs(functions).ParseFiles("./ui/html/base.tmpl")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
ts, err = ts.ParseGlob("./ui/html/partials/nav.tmpl")
|
||||
if err != nil {
|
||||
return nil, err
|
||||
patterns := []string{
|
||||
"html/base.tmpl",
|
||||
"html/partials/*.tmpl",
|
||||
page,
|
||||
}
|
||||
|
||||
ts, err = ts.ParseFiles(page)
|
||||
ts, err := template.New(name).Funcs(functions).ParseFS(ui.Files, patterns...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
39
cmd/web/templates_test.go
Normal file
39
cmd/web/templates_test.go
Normal file
@@ -0,0 +1,39 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gitea.local.lab/Lbenedar/snippetbox/internal/assert"
|
||||
)
|
||||
|
||||
func TestHumanDate(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
tm time.Time
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "UTC",
|
||||
tm: time.Date(2022, 3, 17, 10, 15, 0, 0, time.UTC),
|
||||
want: "17 Mar 2022 at 10:15",
|
||||
},
|
||||
{
|
||||
name: "Empty",
|
||||
tm: time.Time{},
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "CET",
|
||||
tm: time.Date(2022, 3, 17, 10, 15, 0, 0, time.FixedZone("CET", 1*60*60)),
|
||||
want: "17 Mar 2022 at 09:15",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
hd := humanDate(tt.tm)
|
||||
assert.Equal(t, hd, tt.want)
|
||||
})
|
||||
}
|
||||
}
|
||||
104
cmd/web/testutils_test.go
Normal file
104
cmd/web/testutils_test.go
Normal file
@@ -0,0 +1,104 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"html"
|
||||
"io"
|
||||
"log"
|
||||
"net/http"
|
||||
"net/http/cookiejar"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"regexp"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gitea.local.lab/Lbenedar/snippetbox/internal/models/mocks"
|
||||
"github.com/alexedwards/scs/v2"
|
||||
"github.com/go-playground/form/v4"
|
||||
)
|
||||
|
||||
var csrfTokenRX = regexp.MustCompile(`<input type='hidden' name='csrf_token' value='(.+)'>`)
|
||||
|
||||
func extractCSRFToken(t *testing.T, body string) string {
|
||||
t.Logf("Body: %s", body)
|
||||
matches := csrfTokenRX.FindStringSubmatch(body)
|
||||
if len(matches) < 2 {
|
||||
t.Fatal("no csrf token found in body")
|
||||
}
|
||||
|
||||
return html.UnescapeString(string(matches[1]))
|
||||
}
|
||||
|
||||
func newTestApplication(t *testing.T) *application {
|
||||
templateCache, err := newTemplateCache()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
formDecoder := form.NewDecoder()
|
||||
|
||||
sessionManager := scs.New()
|
||||
sessionManager.Lifetime = 12 * time.Hour
|
||||
sessionManager.Cookie.Secure = true
|
||||
|
||||
return &application{
|
||||
errorLog: log.New(io.Discard, "", 0),
|
||||
infoLog: log.New(io.Discard, "", 0),
|
||||
snippets: &mocks.SnippetModel{},
|
||||
users: &mocks.UserModel{},
|
||||
templateCache: templateCache,
|
||||
formDecoder: formDecoder,
|
||||
sessionManager: sessionManager,
|
||||
}
|
||||
}
|
||||
|
||||
type testServer struct {
|
||||
*httptest.Server
|
||||
}
|
||||
|
||||
func newTestServer(t *testing.T, h http.Handler) *testServer {
|
||||
ts := httptest.NewTLSServer(h)
|
||||
jar, err := cookiejar.New(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
ts.Client().Jar = jar
|
||||
|
||||
ts.Client().CheckRedirect = func(req *http.Request, via []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
}
|
||||
return &testServer{ts}
|
||||
}
|
||||
|
||||
func (ts *testServer) get(t *testing.T, urlPath string) (int, http.Header, string) {
|
||||
rs, err := ts.Client().Get(ts.URL + urlPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
defer rs.Body.Close()
|
||||
body, err := io.ReadAll(rs.Body)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
bytes.TrimSpace(body)
|
||||
|
||||
return rs.StatusCode, rs.Header, string(body)
|
||||
}
|
||||
|
||||
func (ts *testServer) postForm(t *testing.T, urlPath string, form url.Values) (int, http.Header, string) {
|
||||
rs, err := ts.Client().PostForm(ts.URL+urlPath, form)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
defer rs.Body.Close()
|
||||
body, err := io.ReadAll(rs.Body)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
bytes.TrimSpace(body)
|
||||
return rs.StatusCode, rs.Header, string(body)
|
||||
}
|
||||
2
go.mod
2
go.mod
@@ -10,4 +10,6 @@ require (
|
||||
github.com/go-sql-driver/mysql v1.9.3 // indirect
|
||||
github.com/julienschmidt/httprouter v1.3.0 // indirect
|
||||
github.com/justinas/alice v1.2.0 // indirect
|
||||
github.com/justinas/nosurf v1.2.0 // indirect
|
||||
golang.org/x/crypto v0.48.0 // indirect
|
||||
)
|
||||
|
||||
4
go.sum
4
go.sum
@@ -13,3 +13,7 @@ github.com/julienschmidt/httprouter v1.3.0 h1:U0609e9tgbseu3rBINet9P48AI/D3oJs4d
|
||||
github.com/julienschmidt/httprouter v1.3.0/go.mod h1:JR6WtHb+2LUe8TCKY3cZOxFyyO8IZAc4RVcycCCAKdM=
|
||||
github.com/justinas/alice v1.2.0 h1:+MHSA/vccVCF4Uq37S42jwlkvI2Xzl7zTPCN5BnZNVo=
|
||||
github.com/justinas/alice v1.2.0/go.mod h1:fN5HRH/reO/zrUflLfTN43t3vXvKzvZIENsNEe7i7qA=
|
||||
github.com/justinas/nosurf v1.2.0 h1:yMs1bSRrNiwXk4AS6n8vL2Ssgpb9CB25T/4xrixaK0s=
|
||||
github.com/justinas/nosurf v1.2.0/go.mod h1:ALpWdSbuNGy2lZWtyXdjkYv4edL23oSEgfBT1gPJ5BQ=
|
||||
golang.org/x/crypto v0.48.0 h1:/VRzVqiRSggnhY7gNRxPauEQ5Drw9haKdM0jqfcCFts=
|
||||
golang.org/x/crypto v0.48.0/go.mod h1:r0kV5h3qnFPlQnBSrULhlsRfryS2pmewsg+XfMgkVos=
|
||||
|
||||
30
internal/assert/assert.go
Normal file
30
internal/assert/assert.go
Normal file
@@ -0,0 +1,30 @@
|
||||
package assert
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func Equal[T comparable](t *testing.T, actual, expected T) {
|
||||
t.Helper()
|
||||
|
||||
if actual != expected {
|
||||
t.Errorf("got: %v; want: %v", actual, expected)
|
||||
}
|
||||
}
|
||||
|
||||
func StringContains(t *testing.T, actual, expectedString string) {
|
||||
t.Helper()
|
||||
|
||||
if !strings.Contains(actual, expectedString) {
|
||||
t.Errorf("got: %q; expected to contain: %q", actual, expectedString)
|
||||
}
|
||||
}
|
||||
|
||||
func NilError(t *testing.T, actual error) {
|
||||
t.Helper()
|
||||
|
||||
if actual != nil {
|
||||
t.Errorf("got: %v; expected: nil", actual)
|
||||
}
|
||||
}
|
||||
34
internal/models/mocks/snippets.go
Normal file
34
internal/models/mocks/snippets.go
Normal file
@@ -0,0 +1,34 @@
|
||||
package mocks
|
||||
|
||||
import (
|
||||
"time"
|
||||
|
||||
"gitea.local.lab/Lbenedar/snippetbox/internal/models"
|
||||
)
|
||||
|
||||
var mockSnippet = &models.Snippet{
|
||||
ID: 1,
|
||||
Title: "An old silent pond",
|
||||
Content: "An old silent pond...",
|
||||
Created: time.Now(),
|
||||
Expires: time.Now(),
|
||||
}
|
||||
|
||||
type SnippetModel struct{}
|
||||
|
||||
func (m *SnippetModel) Insert(title string, content string, expires int) (int, error) {
|
||||
return 2, nil
|
||||
}
|
||||
|
||||
func (m *SnippetModel) Get(id int) (*models.Snippet, error) {
|
||||
switch id {
|
||||
case 1:
|
||||
return mockSnippet, nil
|
||||
default:
|
||||
return nil, models.ErrNoRecord
|
||||
}
|
||||
}
|
||||
|
||||
func (m *SnippetModel) Latest() ([]*models.Snippet, error) {
|
||||
return []*models.Snippet{mockSnippet}, nil
|
||||
}
|
||||
39
internal/models/mocks/users.go
Normal file
39
internal/models/mocks/users.go
Normal file
@@ -0,0 +1,39 @@
|
||||
package mocks
|
||||
|
||||
import "gitea.local.lab/Lbenedar/snippetbox/internal/models"
|
||||
|
||||
type UserModel struct{}
|
||||
|
||||
func (m *UserModel) Insert(name, email, password string) error {
|
||||
switch email {
|
||||
case "dupe@example.com":
|
||||
return models.ErrDuplicateEmail
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func (m *UserModel) Authenticate(email, password string) (int, error) {
|
||||
if email == "alice@example.com" && password == "pa$$word" {
|
||||
return 1, nil
|
||||
}
|
||||
|
||||
return 0, models.ErrInvalidCredentials
|
||||
}
|
||||
|
||||
func (m *UserModel) Exists(id int) (bool, error) {
|
||||
switch id {
|
||||
case 1:
|
||||
return true, nil
|
||||
default:
|
||||
return false, nil
|
||||
}
|
||||
}
|
||||
|
||||
func (m *UserModel) GetById(id int) (*models.User, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (a *UserModel) ChangePassword(id int, curr_pass, new_pass string) error {
|
||||
return nil
|
||||
}
|
||||
@@ -13,7 +13,11 @@ type Snippet struct {
|
||||
Created time.Time
|
||||
Expires time.Time
|
||||
}
|
||||
|
||||
type SnippetModelInterface interface {
|
||||
Insert(title string, content string, expires int) (int, error)
|
||||
Get(id int) (*Snippet, error)
|
||||
Latest() ([]*Snippet, error)
|
||||
}
|
||||
type SnippetModel struct {
|
||||
DB *sql.DB
|
||||
}
|
||||
|
||||
26
internal/models/testdata/setup.sql
vendored
Normal file
26
internal/models/testdata/setup.sql
vendored
Normal file
@@ -0,0 +1,26 @@
|
||||
CREATE TABLE snippets (
|
||||
id INTEGER NOT NULL PRIMARY KEY AUTO_INCREMENT,
|
||||
title VARCHAR(100) NOT NULL,
|
||||
content TEXT NOT NULL,
|
||||
created DATETIME NOT NULL,
|
||||
expires DATETIME NOT NULL
|
||||
);
|
||||
|
||||
CREATE INDEX idx_snippets_created ON snippets(created);
|
||||
|
||||
CREATE TABLE users (
|
||||
id INTEGER NOT NULL PRIMARY KEY AUTO_INCREMENT,
|
||||
name VARCHAR(255) NOT NULL,
|
||||
email VARCHAR(255) NOT NULL,
|
||||
hashed_password CHAR(60) NOT NULL,
|
||||
created DATETIME NOT NULL
|
||||
);
|
||||
|
||||
ALTER TABLE users ADD CONSTRAINT users_uc_email UNIQUE (email);
|
||||
|
||||
INSERT INTO users (name, email, hashed_password, created) VALUES (
|
||||
'Alice Jones',
|
||||
'alice@example.com',
|
||||
'$2a$12$NuTjWXm3KKntReFwyBVHyuf/to.HEwTy.eS206TNfkGfr6HzGJSWG',
|
||||
'2022-01-01 10:00:00'
|
||||
);
|
||||
3
internal/models/testdata/teardown.sql
vendored
Normal file
3
internal/models/testdata/teardown.sql
vendored
Normal file
@@ -0,0 +1,3 @@
|
||||
DROP TABLE users;
|
||||
|
||||
DROP TABLE snippets;
|
||||
37
internal/models/testutils_test.go
Normal file
37
internal/models/testutils_test.go
Normal file
@@ -0,0 +1,37 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"os"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func newTestDB(t *testing.T) *sql.DB {
|
||||
db, err := sql.Open("mysql", "test_web:pass@/test_snippetbox?parseTime=true&multiStatements=true")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
script, err := os.ReadFile("./testdata/setup.sql")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = db.Exec(string(script))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
t.Cleanup(func() {
|
||||
script, err := os.ReadFile("./testdata/teardown.sql")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err = db.Exec(string(script))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
db.Close()
|
||||
})
|
||||
return db
|
||||
}
|
||||
@@ -2,7 +2,12 @@ package models
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"errors"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/go-sql-driver/mysql"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
type User struct {
|
||||
@@ -13,18 +18,112 @@ type User struct {
|
||||
Created time.Time
|
||||
}
|
||||
|
||||
type UserModelInterface interface {
|
||||
Insert(name, email, password string) error
|
||||
Authenticate(email, password string) (int, error)
|
||||
Exists(id int) (bool, error)
|
||||
GetById(id int) (*User, error)
|
||||
ChangePassword(id int, curr_pass, new_pass string) error
|
||||
}
|
||||
|
||||
type UserModel struct {
|
||||
DB *sql.DB
|
||||
}
|
||||
|
||||
func (m *UserModel) Insert(name, email, password string) error {
|
||||
hashedPassword, err := bcrypt.GenerateFromPassword([]byte(password), 12)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
stmt := `INSERT INTO users (name, email, hashed_password, created)
|
||||
VALUES(?, ?, ?, UTC_TIMESTAMP())`
|
||||
|
||||
_, err = m.DB.Exec(stmt, name, email, string(hashedPassword))
|
||||
if err != nil {
|
||||
var mySQLError *mysql.MySQLError
|
||||
if errors.As(err, &mySQLError) {
|
||||
if mySQLError.Number == 1062 && strings.Contains(mySQLError.Message, "users_uc_email") {
|
||||
return ErrDuplicateEmail
|
||||
}
|
||||
}
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *UserModel) Authenticate(email, password string) (int, error) {
|
||||
return 0, nil
|
||||
var id int
|
||||
var hashedPassword []byte
|
||||
|
||||
stmt := `SELECT id, hashed_password FROM users WHERE email = ?`
|
||||
|
||||
err := m.DB.QueryRow(stmt, email).Scan(&id, &hashedPassword)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return 0, ErrInvalidCredentials
|
||||
}
|
||||
return 0, err
|
||||
}
|
||||
err = bcrypt.CompareHashAndPassword(hashedPassword, []byte(password))
|
||||
if err != nil {
|
||||
if errors.Is(err, bcrypt.ErrMismatchedHashAndPassword) {
|
||||
return 0, ErrInvalidCredentials
|
||||
}
|
||||
return 0, err
|
||||
}
|
||||
|
||||
return id, nil
|
||||
}
|
||||
|
||||
func (m *UserModel) Exists(id int) (bool, error) {
|
||||
return false, nil
|
||||
var exists bool
|
||||
|
||||
stmt := `SELECT EXISTS(SELECT true FROM users WHERE id = ?)`
|
||||
|
||||
err := m.DB.QueryRow(stmt, id).Scan(&exists)
|
||||
return exists, err
|
||||
}
|
||||
|
||||
func (a *UserModel) GetById(id int) (*User, error) {
|
||||
acc := User{}
|
||||
|
||||
stmt := `SELECT name, email, created FROM users WHERE id = ?`
|
||||
err := a.DB.QueryRow(stmt, id).Scan(&acc.Name, &acc.Email, &acc.Created)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, ErrInvalidCredentials
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &acc, nil
|
||||
}
|
||||
|
||||
func (a *UserModel) ChangePassword(id int, curr_pass, new_pass string) error {
|
||||
hashedNewPassword, err := bcrypt.GenerateFromPassword([]byte(new_pass), 12)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var hashedDBPassword []byte
|
||||
stmt := `SELECT hashed_password FROM users WHERE id = ?`
|
||||
err = a.DB.QueryRow(stmt, id).Scan(&hashedDBPassword)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
err = bcrypt.CompareHashAndPassword(hashedDBPassword, []byte(curr_pass))
|
||||
if err != nil {
|
||||
if errors.Is(err, bcrypt.ErrMismatchedHashAndPassword) {
|
||||
return ErrInvalidCredentials
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
stmt = `UPDATE users SET hashed_password = ? WHERE id = ?`
|
||||
_, err = a.DB.Exec(stmt, hashedNewPassword, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
45
internal/models/users_test.go
Normal file
45
internal/models/users_test.go
Normal file
@@ -0,0 +1,45 @@
|
||||
package models
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"gitea.local.lab/Lbenedar/snippetbox/internal/assert"
|
||||
)
|
||||
|
||||
func TestUserModelExists(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("models: skipping integration test")
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
userID int
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "Valid ID",
|
||||
userID: 1,
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "Zero ID",
|
||||
userID: 0,
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "Non-existent ID",
|
||||
userID: 2,
|
||||
want: false,
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
db := newTestDB(t)
|
||||
m := UserModel{DB: db}
|
||||
|
||||
exists, err := m.Exists(tt.userID)
|
||||
assert.Equal(t, exists, tt.want)
|
||||
assert.NilError(t, err)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,16 +1,18 @@
|
||||
package validator
|
||||
|
||||
import (
|
||||
"regexp"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
type Validator struct {
|
||||
FieldErrors map[string]string
|
||||
NonFieldErrors []string
|
||||
FieldErrors map[string]string
|
||||
}
|
||||
|
||||
func (v *Validator) Valid() bool {
|
||||
return len(v.FieldErrors) == 0
|
||||
return len(v.FieldErrors) == 0 && len(v.NonFieldErrors) == 0
|
||||
}
|
||||
|
||||
func (v *Validator) AddFieldError(key, message string) {
|
||||
@@ -22,6 +24,10 @@ func (v *Validator) AddFieldError(key, message string) {
|
||||
}
|
||||
}
|
||||
|
||||
func (v *Validator) AddNonFieldError(message string) {
|
||||
v.NonFieldErrors = append(v.NonFieldErrors, message)
|
||||
}
|
||||
|
||||
func (v *Validator) CheckField(ok bool, key, message string) {
|
||||
if !ok {
|
||||
v.AddFieldError(key, message)
|
||||
@@ -36,7 +42,7 @@ func MaxChars(value string, n int) bool {
|
||||
return utf8.RuneCountInString(value) <= n
|
||||
}
|
||||
|
||||
func PermittedInt(value int, permittedValues ...int) bool {
|
||||
func PermittedValue[T comparable](value T, permittedValues ...T) bool {
|
||||
for i := range permittedValues {
|
||||
if value == permittedValues[i] {
|
||||
return true
|
||||
@@ -44,3 +50,13 @@ func PermittedInt(value int, permittedValues ...int) bool {
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
var EmailRX = regexp.MustCompile("^[a-zA-Z0-9.!#$%&'*+\\/=?^_`{|}~-]+@[a-zA-Z0-9](?:[a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?(?:\\.[a-zA-Z0-9](?:[a-zA-Z0-9-]{0,61}[a-zA-Z0-9])?)*$")
|
||||
|
||||
func MinChars(value string, n int) bool {
|
||||
return utf8.RuneCountInString(value) >= n
|
||||
}
|
||||
|
||||
func Matches(value string, rx *regexp.Regexp) bool {
|
||||
return rx.MatchString(value)
|
||||
}
|
||||
|
||||
8
ui/efs.go
Normal file
8
ui/efs.go
Normal file
@@ -0,0 +1,8 @@
|
||||
package ui
|
||||
|
||||
import (
|
||||
"embed"
|
||||
)
|
||||
|
||||
//go:embed "html" "static"
|
||||
var Files embed.FS
|
||||
8
ui/html/pages/about.tmpl
Normal file
8
ui/html/pages/about.tmpl
Normal file
@@ -0,0 +1,8 @@
|
||||
{{define "title"}}About{{end}}
|
||||
|
||||
{{define "main"}}
|
||||
<h2>About</h2>
|
||||
<div class='about'>
|
||||
{{.AboutText}}
|
||||
</div>
|
||||
{{end}}
|
||||
28
ui/html/pages/account.tmpl
Normal file
28
ui/html/pages/account.tmpl
Normal file
@@ -0,0 +1,28 @@
|
||||
{{define "title"}}Account{{end}}
|
||||
|
||||
{{define "main"}}
|
||||
<h2>Account</h2>
|
||||
{{if .Form}}
|
||||
<table>
|
||||
<tr>
|
||||
<th>Name</th>
|
||||
<th>{{.Form.Name}}</th>
|
||||
</tr>
|
||||
<tr>
|
||||
<th>Email</th>
|
||||
<th>{{.Form.Email}}</th>
|
||||
</tr>
|
||||
<tr>
|
||||
<th>Joined</th>
|
||||
<th>{{.Form.Joined}}</th>
|
||||
</tr>
|
||||
</table>
|
||||
{{else}}
|
||||
<p>There's nothing to see here... yet!</p>
|
||||
{{end}}
|
||||
<form action='/account/password' method='GET'>
|
||||
<div>
|
||||
<input type='submit' value='Change password'>
|
||||
</div>
|
||||
</form>
|
||||
{{end}}
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
{{define "main"}}
|
||||
<form action='/snippet/create' method='POST'>
|
||||
<input type='hidden' name='csrf_token' value='{{.CSRFToken}}'>
|
||||
<div>
|
||||
<label>Title:</label>
|
||||
{{with .Form.FieldErrors.title}}
|
||||
|
||||
27
ui/html/pages/login.tmpl
Normal file
27
ui/html/pages/login.tmpl
Normal file
@@ -0,0 +1,27 @@
|
||||
{{define "title"}}Login{{end}}
|
||||
|
||||
{{define "main"}}
|
||||
<form action='/user/login' method='POST' novalidate>
|
||||
<input type='hidden' name='csrf_token' value='{{.CSRFToken}}'>
|
||||
{{range .Form.NonFieldErrors}}
|
||||
<div class='error'>{{.}}</div>
|
||||
{{end}}
|
||||
<div>
|
||||
<label>Email:</label>
|
||||
{{with .Form.FieldErrors.email}}
|
||||
<label class='error'>{{.}}</label>
|
||||
{{end}}
|
||||
<input type='email' name='email' value='{{.Form.Email}}'>
|
||||
</div>
|
||||
<div>
|
||||
<label>Password:</label>
|
||||
{{with .Form.FieldErrors.password}}
|
||||
<label class='error'>{{.}}</label>
|
||||
{{end}}
|
||||
<input type='password' name='password'>
|
||||
</div>
|
||||
<div>
|
||||
<input type='submit' value='Login'>
|
||||
</div>
|
||||
</form>
|
||||
{{end}}
|
||||
35
ui/html/pages/password.tmpl
Normal file
35
ui/html/pages/password.tmpl
Normal file
@@ -0,0 +1,35 @@
|
||||
{{define "title"}}Change Password{{end}}
|
||||
|
||||
{{define "main"}}
|
||||
<h2>Change Password</h2>
|
||||
<form action='/account/password' method='POST' novalidate>
|
||||
<input type='hidden' name='csrf_token' value='{{.CSRFToken}}'>
|
||||
{{range .Form.NonFieldErrors}}
|
||||
<div class='error'>{{.}}</div>
|
||||
{{end}}
|
||||
<div>
|
||||
<label>Current password:</label>
|
||||
{{with .Form.FieldErrors.curr_pass}}
|
||||
<label class='error'>{{.}}</label>
|
||||
{{end}}
|
||||
<input type='password' name='curr_pass' value=''>
|
||||
</div>
|
||||
<div>
|
||||
<label>New password:</label>
|
||||
{{with .Form.FieldErrors.new_pass}}
|
||||
<label class='error'>{{.}}</label>
|
||||
{{end}}
|
||||
<input type='password' name='new_pass'>
|
||||
</div>
|
||||
<div>
|
||||
<label>Confirm new password:</label>
|
||||
{{with .Form.FieldErrors.conf_pass}}
|
||||
<label class='error'>{{.}}</label>
|
||||
{{end}}
|
||||
<input type='password' name='conf_pass'>
|
||||
</div>
|
||||
<div>
|
||||
<input type='submit' value='Change password'>
|
||||
</div>
|
||||
</form>
|
||||
{{end}}
|
||||
31
ui/html/pages/signup.tmpl
Normal file
31
ui/html/pages/signup.tmpl
Normal file
@@ -0,0 +1,31 @@
|
||||
{{define "title"}}Signup{{end}}
|
||||
|
||||
{{define "main"}}
|
||||
<form action='/user/signup' method='POST' novalidate>
|
||||
<input type='hidden' name='csrf_token' value='{{.CSRFToken}}'>
|
||||
<div>
|
||||
<label>Name:</label>
|
||||
{{with .Form.FieldErrors.name}}
|
||||
<label class='error'>{{.}}</label>
|
||||
{{end}}
|
||||
<input type='text' name='name' value='{{.Form.Name}}'>
|
||||
</div>
|
||||
<div>
|
||||
<label>Email:</label>
|
||||
{{with .Form.FieldErrors.email}}
|
||||
<label class='error'>{{.}}</label>
|
||||
{{end}}
|
||||
<input type='email' name='email' value='{{.Form.Email}}'>
|
||||
</div>
|
||||
<div>
|
||||
<label>Password:</label>
|
||||
{{with .Form.FieldErrors.password}}
|
||||
<label class='error'>{{.}}</label>
|
||||
{{end}}
|
||||
<input type='password' name='password'>
|
||||
</div>
|
||||
<div>
|
||||
<input type='submit' value='Signup'>
|
||||
</div>
|
||||
</form>
|
||||
{{end}}
|
||||
@@ -2,14 +2,22 @@
|
||||
<nav>
|
||||
<div>
|
||||
<a href='/'>Home</a>
|
||||
<a href='/snippet/create'>Create snippet</a>
|
||||
<a href='/about'>About</a>
|
||||
{{if .IsAuthenticated}}
|
||||
<a href='/snippet/create'>Create snippet</a>
|
||||
{{end}}
|
||||
</div>
|
||||
<div>
|
||||
<a href='/user/signup'>Signup</a>
|
||||
<a href='/user/login'>Login</a>
|
||||
<form action='/user/logout' method='POST'>
|
||||
<button>Logout</button>
|
||||
</form>
|
||||
{{if .IsAuthenticated}}
|
||||
<a href='/account/view'>Account</a>
|
||||
<form action='/user/logout' method='POST'>
|
||||
<input type='hidden' name='csrf_token' value='{{.CSRFToken}}'>
|
||||
<button>Logout</button>
|
||||
</form>
|
||||
{{else}}
|
||||
<a href='/user/signup'>Signup</a>
|
||||
<a href='/user/login'>Login</a>
|
||||
{{end}}
|
||||
</div>
|
||||
</nav>
|
||||
{{end}}
|
||||
1
ui/static/text/about.txt
Normal file
1
ui/static/text/about.txt
Normal file
@@ -0,0 +1 @@
|
||||
Lorem ipsum dolor sit amet, consectetur adipiscing elit. Phasellus magna quam, eleifend a metus at, tincidunt commodo enim. Integer ac mattis ipsum. Pellentesque habitant morbi tristique senectus et netus et malesuada fames ac turpis egestas. Maecenas non purus ullamcorper, tristique diam vel, sagittis neque. Fusce quam dolor, accumsan non leo sed, auctor rutrum felis. Duis semper tristique tellus, sed sagittis est lobortis nec. Curabitur porttitor lacus eget semper sodales. Duis diam nisl, maximus eu felis id, placerat pulvinar erat. Etiam elementum elit tellus, at auctor ante accumsan vitae. Nulla elit magna, molestie vitae nibh ut, vehicula accumsan ipsum. Ut id consequat ipsum, volutpat volutpat orci. Aenean rhoncus tellus lacus, nec tempor risus interdum sed. Proin quis quam in metus imperdiet lobortis. Duis euismod sem enim, volutpat mollis tortor posuere sed. Morbi condimentum eget dolor at accumsan.
|
||||
Reference in New Issue
Block a user