Files
NianAIGC/backend/internal/postgres/templates.go

137 lines
5.3 KiB
Go

package postgres
import (
"context"
"database/sql"
"encoding/json"
"fmt"
"time"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/templates"
)
const templateColumns = `id, owner_id, name, description, prompt, preview_image_url, settings, sort_order, created_at, updated_at`
const ListImageTemplatesSQL = `SELECT ` + templateColumns + ` FROM public.image_templates WHERE owner_id = $1::text ORDER BY sort_order ASC, updated_at DESC`
const CreateImageTemplateSQL = `INSERT INTO public.image_templates (id, owner_id, name, description, prompt, preview_image_url, settings, sort_order, created_at, updated_at) VALUES ($1::text,$2::text,$3::text,$4::text,$5::text,$6::text,$7::jsonb,$8::integer,$9::timestamptz,$10::timestamptz) RETURNING ` + templateColumns
const UpdateImageTemplateSQL = `UPDATE public.image_templates SET name=COALESCE($3::text,name), description=CASE WHEN $4::boolean THEN $5::text ELSE description END, prompt=COALESCE($6::text,prompt), preview_image_url=CASE WHEN $7::boolean THEN $8::text ELSE preview_image_url END, settings=COALESCE($9::jsonb,settings), sort_order=COALESCE($10::integer,sort_order), updated_at=$11::timestamptz WHERE owner_id=$1::text AND id=$2::text RETURNING ` + templateColumns
const DeleteImageTemplateSQL = `DELETE FROM public.image_templates WHERE owner_id=$1::text AND id=$2::text RETURNING ` + templateColumns
func (db *Database) ListTemplates(ctx context.Context, owner string) ([]templates.Template, error) {
if err := db.available(); err != nil {
return nil, err
}
rows, err := db.querier.Query(ctx, ListImageTemplatesSQL, owner)
if err != nil {
return nil, fmt.Errorf("list image templates: %w", err)
}
defer rows.Close()
items := []templates.Template{}
for rows.Next() {
item, current := scanTemplate(rows)
if current != nil {
return nil, current
}
items = append(items, item)
}
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("read image templates: %w", err)
}
return items, nil
}
func (db *Database) CreateTemplate(ctx context.Context, item templates.Template) (templates.Template, error) {
settings, err := json.Marshal(item.Settings)
if err != nil {
return templates.Template{}, fmt.Errorf("encode template settings: %w", err)
}
return db.requiredTemplate(ctx, CreateImageTemplateSQL, item.ID, item.OwnerID, item.Name, nullableText(item.Description), item.Prompt, nullableText(item.PreviewImageURL), settings, item.SortOrder, item.CreatedAt, item.UpdatedAt)
}
func (db *Database) UpdateTemplate(ctx context.Context, owner, id string, patch templates.Patch, now time.Time) (templates.Template, bool, error) {
var settings any
if patch.Settings != nil {
encoded, err := json.Marshal(*patch.Settings)
if err != nil {
return templates.Template{}, false, fmt.Errorf("encode template settings: %w", err)
}
settings = encoded
}
args := []any{owner, id, patch.Name, patch.Description != nil, pointerText(patch.Description), patch.Prompt, patch.PreviewImageURL != nil, pointerText(patch.PreviewImageURL), settings, patch.SortOrder, now}
return db.oneTemplate(ctx, UpdateImageTemplateSQL, args...)
}
func (db *Database) DeleteTemplate(ctx context.Context, owner, id string) (templates.Template, bool, error) {
return db.oneTemplate(ctx, DeleteImageTemplateSQL, owner, id)
}
func (db *Database) requiredTemplate(ctx context.Context, query string, args ...any) (templates.Template, error) {
item, found, err := db.oneTemplate(ctx, query, args...)
if err != nil {
return templates.Template{}, err
}
if !found {
return templates.Template{}, fmt.Errorf("image template write returned no row")
}
return item, nil
}
func (db *Database) oneTemplate(ctx context.Context, query string, args ...any) (templates.Template, bool, error) {
if err := db.available(); err != nil {
return templates.Template{}, false, err
}
rows, err := db.querier.Query(ctx, query, args...)
if err != nil {
return templates.Template{}, false, fmt.Errorf("query image template: %w", err)
}
defer rows.Close()
if !rows.Next() {
if err := rows.Err(); err != nil {
return templates.Template{}, false, fmt.Errorf("read image template: %w", err)
}
return templates.Template{}, false, nil
}
item, err := scanTemplate(rows)
if err != nil {
return templates.Template{}, false, err
}
if err := rows.Err(); err != nil {
return templates.Template{}, false, fmt.Errorf("read image template: %w", err)
}
return item, true, nil
}
func scanTemplate(rows Rows) (templates.Template, error) {
var item templates.Template
var description, preview sql.NullString
var settings []byte
if err := rows.Scan(&item.ID, &item.OwnerID, &item.Name, &description, &item.Prompt, &preview, &settings, &item.SortOrder, &item.CreatedAt, &item.UpdatedAt); err != nil {
return templates.Template{}, fmt.Errorf("scan image template: %w", err)
}
if description.Valid {
item.Description = description.String
}
if preview.Valid {
item.PreviewImageURL = preview.String
}
if len(settings) == 0 {
settings = []byte(`{}`)
}
if err := json.Unmarshal(settings, &item.Settings); err != nil {
return templates.Template{}, fmt.Errorf("decode template settings: %w", err)
}
return item, nil
}
func nullableText(value string) any {
if value == "" {
return nil
}
return value
}
func pointerText(value *string) any {
if value == nil || *value == "" {
return nil
}
return *value
}
var _ templates.Catalog = (*Database)(nil)