87 lines
2.5 KiB
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
|
|
}
|