Add prompt reference template functions
This commit is contained in:
@@ -4,6 +4,9 @@ import (
|
||||
"bytes"
|
||||
"fmt"
|
||||
"strings"
|
||||
"text/template"
|
||||
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
|
||||
)
|
||||
|
||||
// RenderUserSystem renders the system and user prompt pair for promptID.
|
||||
@@ -16,20 +19,105 @@ func RenderUserSystem(promptID string, data any) (system string, user string, me
|
||||
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 := b.systemTmpl.Execute(&systemBuf, data); err != nil {
|
||||
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 := b.userTmpl.Execute(&userBuf, data); err != nil {
|
||||
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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user