124 lines
4.2 KiB
Go
124 lines
4.2 KiB
Go
package prompt
|
|
|
|
import (
|
|
"bytes"
|
|
"fmt"
|
|
"strings"
|
|
"text/template"
|
|
|
|
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
|
|
)
|
|
|
|
// RenderUserSystem renders the system and user prompt pair for promptID.
|
|
func RenderUserSystem(promptID string, data any) (system string, user string, metadata Metadata, err error) {
|
|
trimmedID := strings.TrimSpace(promptID)
|
|
compiled, ok := promptRegistry[trimmedID]
|
|
if !ok {
|
|
return "", "", Metadata{}, fmt.Errorf("unknown prompt id %q", promptID)
|
|
}
|
|
return compiled.RenderUserSystem(data)
|
|
}
|
|
|
|
// RenderUserSystemWithReferences renders the system and user prompt pair for promptID with reference template functions.
|
|
func RenderUserSystemWithReferences(promptID string, data any, references contracts.ReferenceSet) (system string, user string, metadata Metadata, err error) {
|
|
trimmedID := strings.TrimSpace(promptID)
|
|
compiled, ok := promptRegistry[trimmedID]
|
|
if !ok {
|
|
return "", "", Metadata{}, fmt.Errorf("unknown prompt id %q", promptID)
|
|
}
|
|
return compiled.RenderUserSystemWithReferences(data, references)
|
|
}
|
|
|
|
// RenderUserSystem renders the bundle's system and user prompts.
|
|
func (b *Bundle) RenderUserSystem(data any) (system string, user string, metadata Metadata, err error) {
|
|
return b.RenderUserSystemWithReferences(data, contracts.ReferenceSet{})
|
|
}
|
|
|
|
// RenderUserSystemWithReferences renders the bundle's system and user prompts with reference template functions.
|
|
func (b *Bundle) RenderUserSystemWithReferences(data any, references contracts.ReferenceSet) (system string, user string, metadata Metadata, err error) {
|
|
if b == nil {
|
|
return "", "", Metadata{}, fmt.Errorf("prompt bundle must not be nil")
|
|
}
|
|
systemTmpl, userTmpl, err := b.renderTemplates(references)
|
|
if err != nil {
|
|
return "", "", Metadata{}, err
|
|
}
|
|
var systemBuf bytes.Buffer
|
|
if err := systemTmpl.Execute(&systemBuf, data); err != nil {
|
|
return "", "", Metadata{}, fmt.Errorf("render system prompt %q: %w", b.metadata.PromptID, err)
|
|
}
|
|
|
|
var userBuf bytes.Buffer
|
|
if err := userTmpl.Execute(&userBuf, data); err != nil {
|
|
return "", "", Metadata{}, fmt.Errorf("render user prompt %q: %w", b.metadata.PromptID, err)
|
|
}
|
|
|
|
return strings.TrimSpace(systemBuf.String()), strings.TrimSpace(userBuf.String()), b.metadata, nil
|
|
}
|
|
|
|
func (b *Bundle) renderTemplates(references contracts.ReferenceSet) (*template.Template, *template.Template, error) {
|
|
funcs := b.referenceFuncs(references)
|
|
systemTmpl, err := b.systemTmpl.Clone()
|
|
if err != nil {
|
|
return nil, nil, fmt.Errorf("clone system prompt %q: %w", b.metadata.PromptID, err)
|
|
}
|
|
userTmpl, err := b.userTmpl.Clone()
|
|
if err != nil {
|
|
return nil, nil, fmt.Errorf("clone user prompt %q: %w", b.metadata.PromptID, err)
|
|
}
|
|
systemTmpl.Funcs(funcs)
|
|
userTmpl.Funcs(funcs)
|
|
return systemTmpl, userTmpl, nil
|
|
}
|
|
|
|
func (b *Bundle) referenceFuncs(references contracts.ReferenceSet) template.FuncMap {
|
|
return template.FuncMap{
|
|
"hardening": func() string { return sharedHardening },
|
|
"hasreference": func(slotName string) (bool, error) {
|
|
items, _, err := b.referenceItems(slotName, references)
|
|
if err != nil {
|
|
return false, err
|
|
}
|
|
for _, item := range items {
|
|
if len(item.Content) > 0 {
|
|
return true, nil
|
|
}
|
|
}
|
|
return false, nil
|
|
},
|
|
"reference": func(slotName string) (string, error) {
|
|
items, slot, err := b.referenceItems(slotName, references)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
if len(items) == 0 {
|
|
return "", nil
|
|
}
|
|
if len(items) > 1 && !slot.Multiple {
|
|
return "", fmt.Errorf("reference slot %q has %d bound items but does not allow multiple", slot.Name, len(items))
|
|
}
|
|
parts := make([]string, 0, len(items))
|
|
for _, item := range items {
|
|
parts = append(parts, string(item.Content))
|
|
}
|
|
return strings.Join(parts, "\n"), nil
|
|
},
|
|
}
|
|
}
|
|
|
|
func (b *Bundle) referenceItems(slotName string, references contracts.ReferenceSet) ([]contracts.ReferenceItem, contracts.ReferenceSlot, error) {
|
|
slotName = strings.TrimSpace(slotName)
|
|
slot, ok := b.referenceSlots[slotName]
|
|
if !ok {
|
|
return nil, contracts.ReferenceSlot{}, fmt.Errorf("reference slot %q is not declared", slotName)
|
|
}
|
|
if len(references.Slots) == 0 {
|
|
return nil, slot, nil
|
|
}
|
|
resolved, ok := references.Slots[slotName]
|
|
if !ok {
|
|
return nil, slot, nil
|
|
}
|
|
return append([]contracts.ReferenceItem(nil), resolved.Items...), slot, nil
|
|
}
|