126 lines
3.3 KiB
Go
126 lines
3.3 KiB
Go
package prompt
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"gitea.maximumdirect.net/eric/scriptorium/internal/domain"
|
|
"strings"
|
|
"text/template"
|
|
"unicode/utf8"
|
|
)
|
|
|
|
var (
|
|
ErrMissingRequiredInput = errors.New("missing required input artifact")
|
|
ErrUnknownInput = errors.New("referenced unknown input artifact")
|
|
ErrInvalidTemplate = errors.New("invalid prompt template")
|
|
ErrRenderFailure = errors.New("prompt render failure")
|
|
ErrInvalidMessageRole = errors.New("invalid or empty message role")
|
|
)
|
|
|
|
type goRenderer struct{}
|
|
|
|
func NewGoRenderer() Renderer {
|
|
return &goRenderer{}
|
|
}
|
|
|
|
func (r *goRenderer) Render(ctx context.Context, definition *domain.PromptDefinition, inputs map[string]*domain.Artifact, vars map[string]string) (*domain.RenderedPrompt, error) {
|
|
if definition == nil {
|
|
return nil, fmt.Errorf("%w: nil prompt definition", ErrRenderFailure)
|
|
}
|
|
|
|
// 1. Verify required inputs
|
|
for _, in := range definition.Inputs {
|
|
if !in.Required {
|
|
continue
|
|
}
|
|
art, ok := inputs[in.Name]
|
|
if !ok || art == nil {
|
|
return nil, fmt.Errorf("%w: %s", ErrMissingRequiredInput, in.Name)
|
|
}
|
|
}
|
|
|
|
// 2. Setup template functions
|
|
funcs := template.FuncMap{
|
|
"input": func(name string) (string, error) {
|
|
art, ok := inputs[name]
|
|
if !ok || art == nil {
|
|
return "", fmt.Errorf("%w: %s", ErrUnknownInput, name)
|
|
}
|
|
return string(art.Body), nil
|
|
},
|
|
}
|
|
|
|
sessionID, err := renderSessionID(definition.SessionID, funcs, vars)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
var renderedMessages []domain.RenderedMessage
|
|
|
|
for i, tmplMsg := range definition.Templates {
|
|
select {
|
|
case <-ctx.Done():
|
|
return nil, ctx.Err()
|
|
default:
|
|
}
|
|
|
|
if tmplMsg.Role == "" {
|
|
return nil, fmt.Errorf("%w: message %d", ErrInvalidMessageRole, i)
|
|
}
|
|
|
|
// Parse and execute template
|
|
tmpl, err := template.New(fmt.Sprintf("msg_%d", i)).Funcs(funcs).Option("missingkey=error").Parse(tmplMsg.Content)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("%w: message %d: %v", ErrInvalidTemplate, i, err)
|
|
}
|
|
|
|
var buf bytes.Buffer
|
|
if err := tmpl.Execute(&buf, vars); err != nil {
|
|
return nil, fmt.Errorf("%w: message %d: %w", ErrRenderFailure, i, err)
|
|
}
|
|
|
|
renderedMessages = append(renderedMessages, domain.RenderedMessage{
|
|
Role: tmplMsg.Role,
|
|
Content: buf.String(),
|
|
CacheControl: cloneCacheControl(tmplMsg.CacheControl),
|
|
})
|
|
}
|
|
|
|
return &domain.RenderedPrompt{
|
|
SessionID: sessionID,
|
|
Messages: renderedMessages,
|
|
}, nil
|
|
}
|
|
|
|
func renderSessionID(raw string, funcs template.FuncMap, vars map[string]string) (string, error) {
|
|
if strings.TrimSpace(raw) == "" {
|
|
return "", nil
|
|
}
|
|
|
|
tmpl, err := template.New("session_id").Funcs(funcs).Option("missingkey=error").Parse(raw)
|
|
if err != nil {
|
|
return "", fmt.Errorf("%w: session_id: %v", ErrInvalidTemplate, err)
|
|
}
|
|
|
|
var buf bytes.Buffer
|
|
if err := tmpl.Execute(&buf, vars); err != nil {
|
|
return "", fmt.Errorf("%w: session_id: %w", ErrRenderFailure, err)
|
|
}
|
|
|
|
sessionID := strings.TrimSpace(buf.String())
|
|
if n := utf8.RuneCountInString(sessionID); n > domain.SessionIDMaxLength {
|
|
return "", fmt.Errorf("%w: session_id length %d exceeds maximum %d", ErrRenderFailure, n, domain.SessionIDMaxLength)
|
|
}
|
|
return sessionID, nil
|
|
}
|
|
|
|
func cloneCacheControl(in *domain.CacheControl) *domain.CacheControl {
|
|
if in == nil {
|
|
return nil
|
|
}
|
|
out := *in
|
|
return &out
|
|
}
|