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

87 lines
2.5 KiB
Go

package postgres
import (
"context"
"database/sql"
"fmt"
"testing"
"time"
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/templates"
)
func TestImageTemplateCatalogUsesOwnerScopedExplicitSQL(t *testing.T) {
now := time.Date(2026, 8, 13, 0, 0, 0, 0, time.UTC)
querier := &templateQuerier{rows: &templateRows{rows: [][]any{{"t1", "owner", "Name", nil, "Prompt", nil, []byte(`{"engine":"jimeng"}`), 2, now, now}}}}
db := NewDatabase(Config{Backend: BackendPostgres}, querier)
items, err := db.ListTemplates(context.Background(), "owner")
if err != nil || len(items) != 1 || items[0].Settings.Engine != "jimeng" {
t.Fatalf("items=%#v err=%v", items, err)
}
if querier.query != ListImageTemplatesSQL || len(querier.args) != 1 || querier.args[0] != "owner" {
t.Fatalf("query=%q args=%#v", querier.query, querier.args)
}
}
func TestImageTemplateUpdatePreservesOmittedNullableFields(t *testing.T) {
now := time.Now().UTC()
querier := &templateQuerier{rows: &templateRows{rows: [][]any{{"t1", "owner", "Name", "desc", "Prompt", "/p", []byte(`{}`), 0, now, now}}}}
db := NewDatabase(Config{Backend: BackendPostgres}, querier)
name := "Name"
item, found, err := db.UpdateTemplate(context.Background(), "owner", "t1", templates.Patch{Name: &name}, now)
if err != nil || !found || item.ID != "t1" {
t.Fatalf("item=%#v found=%v err=%v", item, found, err)
}
if querier.query != UpdateImageTemplateSQL || querier.args[3] != false || querier.args[6] != false {
t.Fatalf("unexpected args %#v", querier.args)
}
}
type templateQuerier struct {
rows Rows
query string
args []any
}
func (querier *templateQuerier) Query(_ context.Context, query string, args ...any) (Rows, error) {
querier.query = query
querier.args = args
return querier.rows, nil
}
type templateRows struct {
rows [][]any
index int
}
func (*templateRows) Close() {}
func (*templateRows) Err() error { return nil }
func (rows *templateRows) Next() bool { return rows.index < len(rows.rows) }
func (rows *templateRows) Scan(dest ...any) error {
row := rows.rows[rows.index]
rows.index++
if len(dest) != len(row) {
return fmt.Errorf("scan arity")
}
for index, value := range row {
switch target := dest[index].(type) {
case *string:
if value != nil {
*target = value.(string)
}
case *sql.NullString:
if value != nil {
target.String = value.(string)
target.Valid = true
}
case *[]byte:
*target = value.([]byte)
case *int:
*target = value.(int)
case *time.Time:
*target = value.(time.Time)
}
}
return nil
}