203 lines
4.9 KiB
Go
203 lines
4.9 KiB
Go
package prompt
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"strings"
|
|
"text/template"
|
|
|
|
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
|
|
)
|
|
|
|
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")
|
|
)
|
|
|
|
const artifactTextChunkSize = 64 * 1024
|
|
|
|
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 err := ctx.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
if definition == nil {
|
|
return nil, fmt.Errorf("%w: nil prompt definition", ErrRenderFailure)
|
|
}
|
|
|
|
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)
|
|
}
|
|
}
|
|
|
|
resolver := newArtifactTextResolver(ctx, inputs)
|
|
funcs := template.FuncMap{
|
|
"input": resolver.resolve,
|
|
}
|
|
|
|
if err := ctx.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
sessionID, err := renderSessionID(ctx, definition.SessionID, funcs, vars)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if err := ctx.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
renderedMessages := make([]domain.RenderedMessage, 0, len(definition.Templates))
|
|
|
|
for i, tmplMsg := range definition.Templates {
|
|
if err := ctx.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
if tmplMsg.Role == "" {
|
|
return nil, fmt.Errorf("%w: message %d", ErrInvalidMessageRole, i)
|
|
}
|
|
|
|
if err := ctx.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
tmpl, parseErr := template.New(fmt.Sprintf("msg_%d", i)).Funcs(funcs).Option("missingkey=error").Parse(tmplMsg.Content)
|
|
if err := ctx.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
if parseErr != nil {
|
|
return nil, fmt.Errorf("%w: message %d: %v", ErrInvalidTemplate, i, parseErr)
|
|
}
|
|
|
|
if err := ctx.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
var buf bytes.Buffer
|
|
executeErr := tmpl.Execute(&buf, vars)
|
|
if err := ctx.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
if executeErr != nil {
|
|
return nil, fmt.Errorf("%w: message %d: %w", ErrRenderFailure, i, executeErr)
|
|
}
|
|
|
|
renderedMessages = append(renderedMessages, domain.RenderedMessage{
|
|
Role: tmplMsg.Role,
|
|
Content: buf.String(),
|
|
CacheControl: cloneCacheControl(tmplMsg.CacheControl),
|
|
})
|
|
if err := ctx.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
}
|
|
|
|
if err := ctx.Err(); err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
return &domain.RenderedPrompt{
|
|
SessionID: sessionID,
|
|
Messages: renderedMessages,
|
|
}, nil
|
|
}
|
|
|
|
type artifactTextResolver struct {
|
|
ctx context.Context
|
|
inputs map[string]*domain.Artifact
|
|
textByName map[string]string
|
|
}
|
|
|
|
func newArtifactTextResolver(ctx context.Context, inputs map[string]*domain.Artifact) *artifactTextResolver {
|
|
return &artifactTextResolver{
|
|
ctx: ctx,
|
|
inputs: inputs,
|
|
textByName: make(map[string]string),
|
|
}
|
|
}
|
|
|
|
func (r *artifactTextResolver) resolve(name string) (string, error) {
|
|
if err := r.ctx.Err(); err != nil {
|
|
return "", err
|
|
}
|
|
artifact, ok := r.inputs[name]
|
|
if !ok || artifact == nil {
|
|
return "", fmt.Errorf("%w: %s", ErrUnknownInput, name)
|
|
}
|
|
if text, ok := r.textByName[name]; ok {
|
|
return text, nil
|
|
}
|
|
|
|
var builder strings.Builder
|
|
builder.Grow(len(artifact.Body))
|
|
for start := 0; start < len(artifact.Body); start += artifactTextChunkSize {
|
|
if err := r.ctx.Err(); err != nil {
|
|
return "", err
|
|
}
|
|
end := min(start+artifactTextChunkSize, len(artifact.Body))
|
|
_, _ = builder.Write(artifact.Body[start:end])
|
|
if err := r.ctx.Err(); err != nil {
|
|
return "", err
|
|
}
|
|
}
|
|
if err := r.ctx.Err(); err != nil {
|
|
return "", err
|
|
}
|
|
|
|
text := builder.String()
|
|
r.textByName[name] = text
|
|
return text, nil
|
|
}
|
|
|
|
func renderSessionID(ctx context.Context, raw string, funcs template.FuncMap, vars map[string]string) (string, error) {
|
|
if err := ctx.Err(); err != nil {
|
|
return "", err
|
|
}
|
|
tmpl, parseErr := template.New("session_id").Funcs(funcs).Option("missingkey=error").Parse(raw)
|
|
if err := ctx.Err(); err != nil {
|
|
return "", err
|
|
}
|
|
if parseErr != nil {
|
|
return "", fmt.Errorf("%w: session_id: %v", ErrInvalidTemplate, parseErr)
|
|
}
|
|
|
|
if err := ctx.Err(); err != nil {
|
|
return "", err
|
|
}
|
|
var buf bytes.Buffer
|
|
executeErr := tmpl.Execute(&buf, vars)
|
|
if err := ctx.Err(); err != nil {
|
|
return "", err
|
|
}
|
|
if executeErr != nil {
|
|
return "", fmt.Errorf("%w: session_id: %w", ErrRenderFailure, executeErr)
|
|
}
|
|
|
|
sessionID, err := domain.NormalizeSessionID(buf.String())
|
|
if err != nil {
|
|
return "", fmt.Errorf("%w: session_id: %v", ErrRenderFailure, err)
|
|
}
|
|
return sessionID, nil
|
|
}
|
|
|
|
func cloneCacheControl(in *domain.CacheControl) *domain.CacheControl {
|
|
if in == nil {
|
|
return nil
|
|
}
|
|
out := *in
|
|
return &out
|
|
}
|