Files
chorus/internal/core/generate/prompt_test.go
T

85 lines
2.7 KiB
Go

package generate
import (
"errors"
"testing"
"git.ilapage.cn/OPC/chorus/internal/core/model"
)
func TestRenderPromptGolden(t *testing.T) {
tests := []struct {
name string
template string
prompt string
want string
}{
{
name: "chat preserves user text",
template: "{{.UserPrompt}}",
prompt: " 保留前后空格 ",
want: " 保留前后空格 ",
},
{
name: "images edits confirmed template",
template: "用户要求:\n{{.UserPrompt}}\n\n图片说明:\n- 第 1 张图片是需要处理的主图。\n- 其余图片仅作为参考。\n- 仅按用户明确要求进行修改,不添加用户未要求的风格、文字、商品属性或场景。",
prompt: "把背景改为白色",
want: "用户要求:\n把背景改为白色\n\n图片说明:\n- 第 1 张图片是需要处理的主图。\n- 其余图片仅作为参考。\n- 仅按用户明确要求进行修改,不添加用户未要求的风格、文字、商品属性或场景。",
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
got, err := RenderPrompt(test.template, test.prompt)
if err != nil {
t.Fatalf("RenderPrompt() error = %v", err)
}
if got != test.want {
t.Fatalf("RenderPrompt() = %q, want %q", got, test.want)
}
})
}
}
func TestRenderPromptRejectsUnsafeOrIncompleteTemplates(t *testing.T) {
tests := []struct {
name string
template string
prompt string
want error
}{
{"empty input", "{{.UserPrompt}}", " ", ErrEmptyUserPrompt},
{"parse failure", "{{.UserPrompt", "hello", ErrInvalidTemplate},
{"missing variable", "static", "hello", ErrInvalidTemplate},
{"unknown variable", "{{.Other}}", "hello", ErrInvalidTemplate},
{"function call", "{{printf `%s` .UserPrompt}}", "hello", ErrInvalidTemplate},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
_, err := RenderPrompt(test.template, test.prompt)
if !errors.Is(err, test.want) {
t.Fatalf("RenderPrompt() error = %v, want %v", err, test.want)
}
})
}
}
func TestAssignImageRolesPreservesOrder(t *testing.T) {
inputs := []model.GenerationInput{{OriginalName: "a.png"}, {OriginalName: "b.png"}, {OriginalName: "c.png"}}
assigned, err := AssignImageRoles(inputs)
if err != nil {
t.Fatalf("AssignImageRoles() error = %v", err)
}
for index, input := range assigned {
if input.OriginalName != inputs[index].OriginalName || input.Position != uint32(index) {
t.Fatalf("input %d order changed: %#v", index, input)
}
wantRole := model.RoleReference
if index == 0 {
wantRole = model.RolePrimary
}
if input.Role != wantRole {
t.Errorf("input %d role = %s, want %s", index, input.Role, wantRole)
}
}
}