package main
import (
"context"
"embed"
"encoding/json"
"errors"
"io"
"net/http"
"strconv"
"strings"
"github.com/jackc/pgx/v5/pgconn"
)
//go:embed compat_errors/*.txt openapi.json docs.html
var compatibilityErrors embed.FS
type Application struct {
store *Store
production bool
}
func writeJSON(w http.ResponseWriter, status int, v any) {
w.Header().Set("Content-Type", "application/json; charset=utf-8")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(v)
}
func (a *Application) failure(w http.ResponseWriter, err error) {
var miss missing
if errors.As(err, &miss) {
writeJSON(w, 404, map[string]string{"detail": err.Error()})
return
}
if a.production {
writeJSON(w, 500, map[string]string{"detail": "Internal Server Error"})
return
}
name := ""
var dupe duplicateUsers
var tag missingTag
var pg *pgconn.PgError
if errors.As(err, &dupe) {
name = "email_duplicate"
} else if errors.As(err, &tag) {
name = "create_unknown_tag"
} else if errors.Is(err, errTitleTooLong) || (errors.As(err, &pg) && pg.Code == "22001") {
name = "create_long_title"
}
if name != "" {
raw, _ := compatibilityErrors.ReadFile("compat_errors/" + name + ".txt")
text := string(raw)
if dupe > 0 {
count := strconv.Itoa(int(dupe))
if dupe == 21 {
count = "more than 20"
}
text = strings.ReplaceAll(text, "it returned 2!", "it returned "+count+"!")
}
w.Header().Set("Content-Type", "text/plain")
w.WriteHeader(500)
io.WriteString(w, text)
return
}
writeJSON(w, 500, map[string]string{"detail": "Internal Server Error"})
}
func method(w http.ResponseWriter, r *http.Request, allow string) bool {
for _, m := range strings.Split(allow, ", ") {
if m == r.Method {
return true
}
}
w.Header().Set("Content-Type", "text/html; charset=utf-8")
w.Header().Set("Allow", allow)
w.WriteHeader(405)
if r.Method != "HEAD" {
io.WriteString(w, "Method not allowed")
}
return false
}
func pathID(w http.ResponseWriter, value, name string) (int64, bool) {
id, kind := integer(value)
if kind == "overflow" {
model := "Post"
if name == "user_id" {
model = "User"
}
writeJSON(w, 404, map[string]string{"detail": missing(model).Error()})
return 0, false
}
if kind != "" {
writeJSON(w, 422, map[string]any{"detail": []Validation{issue(kind, "path", name)}})
return 0, false
}
return id, true
}
func argument(w http.ResponseWriter, r *http.Request, name string) (string, bool) {
v, exists := r.URL.Query()[name]
if !exists {
writeJSON(w, 422, map[string]any{"detail": []Validation{issue("missing", "query", name)}})
return "", false
}
return v[len(v)-1], true
}
func (a *Application) ServeHTTP(w http.ResponseWriter, r *http.Request) {
w.Header().Set("X-Frame-Options", "DENY")
w.Header().Set("X-Content-Type-Options", "nosniff")
w.Header().Set("Referrer-Policy", "same-origin")
w.Header().Set("Cross-Origin-Opener-Policy", "same-origin")
path := r.URL.Path
ctx := r.Context()
switch {
case path == "/api/openapi.json" || path == "/api/docs":
if !method(w, r, "GET") {
return
}
name := "openapi.json"
contentType := "application/json"
if path == "/api/docs" {
name = "docs.html"
contentType = "text/html; charset=utf-8"
}
raw, _ := compatibilityErrors.ReadFile(name)
w.Header().Set("Content-Type", contentType)
w.Write(raw)
case path == "/api/posts":
if !method(w, r, "POST, GET") {
return
}
if r.Method == "GET" {
a.list(w, ctx, "", nil)
return
}
a.create(w, r, "", true)
case path == "/api/posts/search":
if !method(w, r, "GET") {
return
}
q, ok := argument(w, r, "q")
if ok {
a.list(w, ctx, "search", q)
}
case strings.HasPrefix(path, "/api/posts/by-tag/") && len(strings.Split(path, "/")) == 5:
if !method(w, r, "GET") {
return
}
a.list(w, ctx, "tag", strings.TrimPrefix(path, "/api/posts/by-tag/"))
case strings.HasPrefix(path, "/api/posts/"):
parts := strings.Split(path, "/")
if len(parts) == 4 {
if !method(w, r, "GET") {
return
}
id, ok := pathID(w, parts[3], "post_id")
if !ok {
return
}
data, err := a.store.detail(ctx, id)
if err != nil {
a.failure(w, err)
} else {
writeJSON(w, 200, data)
}
} else if len(parts) == 5 && parts[4] == "comments" {
if !method(w, r, "POST") {
return
}
a.create(w, r, parts[3], false)
} else {
http.NotFound(w, r)
}
case path == "/api/users/find":
if !method(w, r, "GET") {
return
}
email, ok := argument(w, r, "email")
if !ok {
return
}
data, err := a.store.user(ctx, "email", email)
if err != nil {
a.failure(w, err)
} else {
writeJSON(w, 200, data)
}
case strings.HasPrefix(path, "/api/users/") && len(strings.Split(path, "/")) == 4:
if !method(w, r, "GET") {
return
}
id, ok := pathID(w, strings.TrimPrefix(path, "/api/users/"), "user_id")
if !ok {
return
}
data, err := a.store.user(ctx, "id", id)
if err != nil {
a.failure(w, err)
} else {
writeJSON(w, 200, data)
}
default:
http.NotFound(w, r)
}
}
func (a *Application) list(w http.ResponseWriter, ctx context.Context, filter string, arg any) {
posts, err := a.store.list(ctx, filter, arg)
if err != nil {
a.failure(w, err)
} else {
writeJSON(w, 200, posts)
}
}
func (a *Application) create(w http.ResponseWriter, r *http.Request, postRaw string, post bool) {
raw, err := io.ReadAll(r.Body)
if err != nil {
var oversized *http.MaxBytesError
if a.production && errors.As(err, &oversized) {
writeJSON(w, 413, map[string]string{"detail": "Request body too large"})
return
}
writeJSON(w, 400, map[string]string{"detail": "Cannot read request body"})
return
}
input, issues, err := validate(raw, post)
if err != nil {
writeJSON(w, 400, map[string]string{"detail": "Cannot parse request body (" + err.Error() + ")"})
return
}
var postID int64
kind := ""
if !post {
postID, kind = integer(postRaw)
if kind != "" && kind != "overflow" {
issues = append([]Validation{issue(kind, "path", "post_id")}, issues...)
}
}
if len(issues) > 0 {
writeJSON(w, 422, map[string]any{"detail": issues})
return
}
if kind == "overflow" {
a.failure(w, missing("Post"))
return
}
var id int64
if post {
id, err = a.store.create(r.Context(), input)
} else {
id, err = a.store.comment(r.Context(), postID, input)
}
if err != nil {
a.failure(w, err)
return
}
if post {
writeJSON(w, 200, map[string]any{"id": id, "title": input.Title})
} else {
writeJSON(w, 200, map[string]any{"id": id})
}
}