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 }