74 lines
2.0 KiB
Go
74 lines
2.0 KiB
Go
package generate
|
|
|
|
import (
|
|
"bytes"
|
|
"errors"
|
|
"fmt"
|
|
"strings"
|
|
"text/template"
|
|
"text/template/parse"
|
|
)
|
|
|
|
var (
|
|
ErrEmptyUserPrompt = errors.New("user prompt is empty")
|
|
ErrInvalidTemplate = errors.New("prompt template is invalid")
|
|
)
|
|
|
|
type PromptData struct {
|
|
UserPrompt string
|
|
}
|
|
|
|
func RenderPrompt(templateText, userPrompt string) (string, error) {
|
|
if strings.TrimSpace(userPrompt) == "" {
|
|
return "", ErrEmptyUserPrompt
|
|
}
|
|
tmpl, err := template.New("prompt").Option("missingkey=error").Parse(templateText)
|
|
if err != nil {
|
|
return "", fmt.Errorf("%w: parse: %v", ErrInvalidTemplate, err)
|
|
}
|
|
if len(tmpl.Templates()) != 1 {
|
|
return "", fmt.Errorf("%w: nested templates are not allowed", ErrInvalidTemplate)
|
|
}
|
|
found, err := validatePromptNode(tmpl.Tree.Root)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
if !found {
|
|
return "", fmt.Errorf("%w: required variable .UserPrompt is missing", ErrInvalidTemplate)
|
|
}
|
|
|
|
var rendered bytes.Buffer
|
|
if err := tmpl.Execute(&rendered, PromptData{UserPrompt: userPrompt}); err != nil {
|
|
return "", fmt.Errorf("%w: execute: %v", ErrInvalidTemplate, err)
|
|
}
|
|
return rendered.String(), nil
|
|
}
|
|
|
|
func validatePromptNode(node parse.Node) (bool, error) {
|
|
switch typed := node.(type) {
|
|
case *parse.ListNode:
|
|
found := false
|
|
for _, child := range typed.Nodes {
|
|
childFound, err := validatePromptNode(child)
|
|
if err != nil {
|
|
return false, err
|
|
}
|
|
found = found || childFound
|
|
}
|
|
return found, nil
|
|
case *parse.TextNode, *parse.CommentNode:
|
|
return false, nil
|
|
case *parse.ActionNode:
|
|
if len(typed.Pipe.Decl) != 0 || len(typed.Pipe.Cmds) != 1 || len(typed.Pipe.Cmds[0].Args) != 1 {
|
|
return false, fmt.Errorf("%w: only direct .UserPrompt interpolation is allowed", ErrInvalidTemplate)
|
|
}
|
|
field, ok := typed.Pipe.Cmds[0].Args[0].(*parse.FieldNode)
|
|
if !ok || len(field.Ident) != 1 || field.Ident[0] != "UserPrompt" {
|
|
return false, fmt.Errorf("%w: only .UserPrompt is allowed", ErrInvalidTemplate)
|
|
}
|
|
return true, nil
|
|
default:
|
|
return false, fmt.Errorf("%w: unsupported template construct %T", ErrInvalidTemplate, node)
|
|
}
|
|
}
|