diff --git a/cmd/web/handlers.go b/cmd/web/handlers.go index 8605a1b..7fbf19e 100644 --- a/cmd/web/handlers.go +++ b/cmd/web/handlers.go @@ -134,8 +134,9 @@ func (app *application) snippetCreatePost(w http.ResponseWriter, r *http.Request func (app *application) userSignupPost(w http.ResponseWriter, r *http.Request) { var form userSignupForm err := app.decodePostForm(r, &form) + if err != nil { - app.clientError(w, http.StatusBadRequest) + app.notFound(w) return } diff --git a/cmd/web/handlers_test.go b/cmd/web/handlers_test.go index c3dc007..42574e7 100644 --- a/cmd/web/handlers_test.go +++ b/cmd/web/handlers_test.go @@ -2,6 +2,7 @@ package main import ( "net/http" + "net/url" "testing" "gitea.local.lab/Lbenedar/snippetbox/internal/assert" @@ -76,3 +77,118 @@ func TestSnippetView(t *testing.T) { }) } } + +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 = "
" + ) + + 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) + } + }) + } +} diff --git a/cmd/web/testutils_test.go b/cmd/web/testutils_test.go index 01cdc88..6b6e6aa 100644 --- a/cmd/web/testutils_test.go +++ b/cmd/web/testutils_test.go @@ -2,11 +2,14 @@ package main import ( "bytes" + "html" "io" "log" "net/http" "net/http/cookiejar" "net/http/httptest" + "net/url" + "regexp" "testing" "time" @@ -15,6 +18,18 @@ import ( "github.com/go-playground/form/v4" ) +var csrfTokenRX = regexp.MustCompile(``) + +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 { @@ -72,3 +87,18 @@ func (ts *testServer) get(t *testing.T, urlPath string) (int, http.Header, strin 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) +} diff --git a/internal/assert/assert.go b/internal/assert/assert.go index ee89af1..19e86a8 100644 --- a/internal/assert/assert.go +++ b/internal/assert/assert.go @@ -20,3 +20,11 @@ func StringContains(t *testing.T, actual, expectedString string) { 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) + } +} diff --git a/internal/models/testdata/setup.sql b/internal/models/testdata/setup.sql new file mode 100644 index 0000000..3dc9c29 --- /dev/null +++ b/internal/models/testdata/setup.sql @@ -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' +); \ No newline at end of file diff --git a/internal/models/testdata/teardown.sql b/internal/models/testdata/teardown.sql new file mode 100644 index 0000000..7f84322 --- /dev/null +++ b/internal/models/testdata/teardown.sql @@ -0,0 +1,3 @@ +DROP TABLE users; + +DROP TABLE snippets; \ No newline at end of file diff --git a/internal/models/testutils_test.go b/internal/models/testutils_test.go new file mode 100644 index 0000000..920695b --- /dev/null +++ b/internal/models/testutils_test.go @@ -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 +} diff --git a/internal/models/users_test.go b/internal/models/users_test.go new file mode 100644 index 0000000..332f86f --- /dev/null +++ b/internal/models/users_test.go @@ -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) + }) + } +}