137 lines
5.3 KiB
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)
|