package main
import (
"context"
"errors"
"fmt"
"strings"
"time"
"unicode/utf8"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
)
type Author struct {
ID int64 `json:"id"`
Username string `json:"username"`
DisplayName string `json:"display_name"`
}
type Tag struct {
ID int64 `json:"id"`
Name string `json:"name"`
Slug string `json:"slug"`
}
type Post struct {
ID int64 `json:"id"`
Title string `json:"title"`
Author Author `json:"author"`
Tags []Tag `json:"tags"`
ViewCount int `json:"view_count"`
CreatedAt string `json:"created_at"`
}
type Comment struct {
ID int64 `json:"id"`
Author Author `json:"author"`
Body string `json:"body"`
CreatedAt string `json:"created_at"`
}
type Detail struct {
Post
Body string `json:"body"`
Comments []Comment `json:"comments"`
UpdatedAt string `json:"updated_at"`
}
type User struct {
Author
Email string `json:"email"`
Bio string `json:"bio"`
PostCount int64 `json:"post_count"`
CommentCount int64 `json:"comment_count"`
}
type Store struct{ pool *pgxpool.Pool }
type missing string
func (m missing) Error() string { return "Not Found: No " + string(m) + " matches the given query." }
type duplicateUsers int
func (d duplicateUsers) Error() string { return fmt.Sprintf("duplicate email: %d users", d) }
type missingTag struct{}
var errTitleTooLong = errors.New("value too long for type character varying(255)")
func (missingTag) Error() string { return "Tag matching query does not exist." }
func date(t time.Time) string {
t = t.UTC()
if t.Nanosecond() == 0 {
return t.Format("2006-01-02T15:04:05Z")
}
return t.Format("2006-01-02T15:04:05.000Z")
}
func like(q string) string {
return "%" + strings.NewReplacer(`\`, `\\`, `%`, `\%`, `_`, `\_`).Replace(q) + "%"
}
func (s *Store) list(ctx context.Context, filter string, arg any) ([]Post, error) {
predicate := "p.is_published"
args := []any{}
switch filter {
case "search":
predicate += " AND (UPPER(p.title::text) LIKE UPPER($1) OR UPPER(p.body::text) LIKE UPPER($1))"
args = append(args, like(arg.(string)))
case "tag":
var id int64
if err := s.pool.QueryRow(ctx, "SELECT id FROM blog_tag WHERE slug=$1", arg).Scan(&id); err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return nil, missing("Tag")
}
return nil, err
}
predicate += " AND EXISTS (SELECT 1 FROM blog_post_tags pt WHERE pt.post_id=p.id AND pt.tag_id=$1)"
args = append(args, id)
}
rows, err := s.pool.Query(ctx, `SELECT p.id,p.title,p.view_count,p.created_at,u.id,u.username,u.display_name FROM blog_post p JOIN blog_user u ON u.id=p.author_id WHERE `+predicate+` ORDER BY p.created_at DESC,p.id`, args...)
if err != nil {
return nil, err
}
posts := []Post{}
ids := []int64{}
positions := map[int64]int{}
for rows.Next() {
p := Post{Tags: []Tag{}}
var ts time.Time
err = rows.Scan(&p.ID, &p.Title, &p.ViewCount, &ts, &p.Author.ID, &p.Author.Username, &p.Author.DisplayName)
if err != nil {
rows.Close()
return nil, err
}
p.CreatedAt = date(ts)
positions[p.ID] = len(posts)
posts = append(posts, p)
ids = append(ids, p.ID)
}
err = rows.Err()
rows.Close()
if err != nil {
return nil, err
}
if len(posts) == 0 {
return posts, nil
}
rows, err = s.pool.Query(ctx, `SELECT pt.post_id,t.id,t.name,t.slug FROM blog_post_tags pt JOIN blog_tag t ON t.id=pt.tag_id WHERE pt.post_id=ANY($1::bigint[])`, ids)
if err != nil {
return nil, err
}
defer rows.Close()
for rows.Next() {
var pid int64
var tag Tag
if err = rows.Scan(&pid, &tag.ID, &tag.Name, &tag.Slug); err != nil {
return nil, err
}
i := positions[pid]
posts[i].Tags = append(posts[i].Tags, tag)
}
return posts, rows.Err()
}
func (s *Store) detail(ctx context.Context, id int64) (Detail, error) {
p := Detail{Post: Post{Tags: []Tag{}}, Comments: []Comment{}}
var created time.Time
var published bool
err := s.pool.QueryRow(ctx, `SELECT p.id,p.title,p.body,p.is_published,p.view_count,p.created_at,u.id,u.username,u.display_name FROM blog_post p JOIN blog_user u ON u.id=p.author_id WHERE p.id=$1`, id).Scan(&p.ID, &p.Title, &p.Body, &published, &p.ViewCount, &created, &p.Author.ID, &p.Author.Username, &p.Author.DisplayName)
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return p, missing("Post")
}
return p, err
}
p.ViewCount++
updated := time.Now().UTC().Truncate(time.Microsecond)
p.CreatedAt = date(created)
p.UpdatedAt = date(updated)
// Deliberately retain Django's read/modify/full-row-save behavior and race.
_, err = s.pool.Exec(ctx, `UPDATE blog_post SET title=$1,body=$2,is_published=$3,view_count=$4,created_at=$5,updated_at=$6,author_id=$7 WHERE id=$8`, p.Title, p.Body, published, p.ViewCount, created, updated, p.Author.ID, id)
if err != nil {
return p, err
}
rows, err := s.pool.Query(ctx, `SELECT t.id,t.name,t.slug FROM blog_tag t JOIN blog_post_tags pt ON pt.tag_id=t.id WHERE pt.post_id=$1`, id)
if err != nil {
return p, err
}
for rows.Next() {
var tag Tag
if err = rows.Scan(&tag.ID, &tag.Name, &tag.Slug); err != nil {
rows.Close()
return p, err
}
p.Tags = append(p.Tags, tag)
}
err = rows.Err()
rows.Close()
if err != nil {
return p, err
}
rows, err = s.pool.Query(ctx, `SELECT c.id,c.body,c.created_at,u.id,u.username,u.display_name FROM blog_comment c JOIN blog_user u ON u.id=c.author_id WHERE c.post_id=$1 ORDER BY c.created_at,c.id`, id)
if err != nil {
return p, err
}
defer rows.Close()
for rows.Next() {
var comment Comment
var ts time.Time
if err = rows.Scan(&comment.ID, &comment.Body, &ts, &comment.Author.ID, &comment.Author.Username, &comment.Author.DisplayName); err != nil {
return p, err
}
comment.CreatedAt = date(ts)
p.Comments = append(p.Comments, comment)
}
return p, rows.Err()
}
func (s *Store) user(ctx context.Context, by string, value any) (User, error) {
where := "u.id=$1"
if by == "email" {
where = "u.email=$1"
}
rows, err := s.pool.Query(ctx, `SELECT u.id,u.username,u.display_name,u.email,u.bio,(SELECT count(*) FROM blog_post p WHERE p.author_id=u.id),(SELECT count(*) FROM blog_comment c WHERE c.author_id=u.id) FROM blog_user u WHERE `+where+` LIMIT 21`, value)
if err != nil {
return User{}, err
}
defer rows.Close()
var user User
count := 0
for rows.Next() {
count++
if err = rows.Scan(&user.ID, &user.Username, &user.DisplayName, &user.Email, &user.Bio, &user.PostCount, &user.CommentCount); err != nil {
return User{}, err
}
}
if err = rows.Err(); err != nil {
return User{}, err
}
if count == 0 {
return User{}, missing("User")
}
if count > 1 {
return User{}, duplicateUsers(count)
}
return user, nil
}
func (s *Store) exists(ctx context.Context, table, model string, id int64) error {
var found int64
err := s.pool.QueryRow(ctx, "SELECT id FROM "+table+" WHERE id=$1", id).Scan(&found)
if errors.Is(err, pgx.ErrNoRows) {
return missing(model)
}
return err
}
func (s *Store) create(ctx context.Context, input Input) (int64, error) {
if input.AuthorOverflow {
return 0, missing("User")
}
if err := s.exists(ctx, "blog_user", "User", input.AuthorID); err != nil {
return 0, err
}
// psycopg's unknown-type bind fails varchar coercion before nextval; pgx's
// prepared insert otherwise consumes an ID. Retain the 500 and sequence state.
if utf8.RuneCountInString(input.Title) > 255 {
return 0, errTitleTooLong
}
var id int64
now := time.Now().UTC().Truncate(time.Microsecond)
err := s.pool.QueryRow(ctx, `INSERT INTO blog_post(title,body,is_published,view_count,created_at,updated_at,author_id) VALUES($1,$2,true,0,$3,$3,$4) RETURNING id`, input.Title, input.Body, now, input.AuthorID).Scan(&id)
if err != nil {
return 0, err
}
// No wrapping transaction: unknown later tags intentionally leave partial writes.
for _, slug := range input.Tags {
var tid int64
if err = s.pool.QueryRow(ctx, `SELECT id FROM blog_tag WHERE slug=$1`, slug).Scan(&tid); err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return id, missingTag{}
}
return id, err
}
if _, err = s.pool.Exec(ctx, `INSERT INTO blog_post_tags(post_id,tag_id) VALUES($1,$2) ON CONFLICT DO NOTHING`, id, tid); err != nil {
return id, err
}
}
return id, nil
}
func (s *Store) comment(ctx context.Context, post int64, input Input) (int64, error) {
if err := s.exists(ctx, "blog_post", "Post", post); err != nil {
return 0, err
}
if input.AuthorOverflow {
return 0, missing("User")
}
if err := s.exists(ctx, "blog_user", "User", input.AuthorID); err != nil {
return 0, err
}
var id int64
err := s.pool.QueryRow(ctx, `INSERT INTO blog_comment(body,created_at,post_id,author_id) VALUES($1,$2,$3,$4) RETURNING id`, input.Body, time.Now().UTC().Truncate(time.Microsecond), post, input.AuthorID).Scan(&id)
return id, err
}