Add prompt reference template functions
This commit is contained in:
@@ -7,9 +7,13 @@ import (
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"path"
|
||||
"reflect"
|
||||
"sort"
|
||||
"strings"
|
||||
"text/template"
|
||||
"text/template/parse"
|
||||
|
||||
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
|
||||
)
|
||||
|
||||
//go:embed assets/**
|
||||
@@ -43,18 +47,20 @@ func (m Metadata) DiagnosticsMap() map[string]any {
|
||||
|
||||
// Definition identifies a caller-owned system/user prompt bundle.
|
||||
type Definition struct {
|
||||
PromptID string
|
||||
Version string
|
||||
EmbeddedPath string
|
||||
SystemPath string
|
||||
UserPath string
|
||||
PromptID string
|
||||
Version string
|
||||
EmbeddedPath string
|
||||
SystemPath string
|
||||
UserPath string
|
||||
ReferenceSlots []contracts.ReferenceSlot
|
||||
}
|
||||
|
||||
// Bundle is a compiled system/user prompt pair.
|
||||
type Bundle struct {
|
||||
systemTmpl *template.Template
|
||||
userTmpl *template.Template
|
||||
metadata Metadata
|
||||
systemTmpl *template.Template
|
||||
userTmpl *template.Template
|
||||
metadata Metadata
|
||||
referenceSlots map[string]contracts.ReferenceSlot
|
||||
}
|
||||
|
||||
// Metadata returns metadata for the compiled prompt bundle.
|
||||
@@ -168,7 +174,9 @@ func LoadBundle(fsys fs.FS, def Definition) (*Bundle, error) {
|
||||
}
|
||||
|
||||
funcs := template.FuncMap{
|
||||
"hardening": func() string { return sharedHardening },
|
||||
"hardening": func() string { return sharedHardening },
|
||||
"reference": func(string) (string, error) { return "", nil },
|
||||
"hasreference": func(string) (bool, error) { return false, nil },
|
||||
}
|
||||
systemTmpl, err := template.New(path.Base(systemPath)).Option("missingkey=error").Funcs(funcs).Parse(systemSource)
|
||||
if err != nil {
|
||||
@@ -178,6 +186,13 @@ func LoadBundle(fsys fs.FS, def Definition) (*Bundle, error) {
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parse embedded user prompt %q: %w", userPath, err)
|
||||
}
|
||||
referenceSlots := referenceSlotMap(def.ReferenceSlots)
|
||||
if err := validateTemplateReferenceSlots(systemTmpl, referenceSlots); err != nil {
|
||||
return nil, fmt.Errorf("validate embedded system prompt %q: %w", systemPath, err)
|
||||
}
|
||||
if err := validateTemplateReferenceSlots(userTmpl, referenceSlots); err != nil {
|
||||
return nil, fmt.Errorf("validate embedded user prompt %q: %w", userPath, err)
|
||||
}
|
||||
|
||||
hashInput := systemSource + "\n\n" + userSource
|
||||
hash := sha256.Sum256([]byte(hashInput))
|
||||
@@ -190,12 +205,127 @@ func LoadBundle(fsys fs.FS, def Definition) (*Bundle, error) {
|
||||
}
|
||||
|
||||
return &Bundle{
|
||||
systemTmpl: systemTmpl,
|
||||
userTmpl: userTmpl,
|
||||
metadata: metadata,
|
||||
systemTmpl: systemTmpl,
|
||||
userTmpl: userTmpl,
|
||||
metadata: metadata,
|
||||
referenceSlots: referenceSlots,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func referenceSlotMap(slots []contracts.ReferenceSlot) map[string]contracts.ReferenceSlot {
|
||||
if len(slots) == 0 {
|
||||
return nil
|
||||
}
|
||||
out := make(map[string]contracts.ReferenceSlot, len(slots))
|
||||
for _, slot := range slots {
|
||||
name := strings.TrimSpace(slot.Name)
|
||||
if name == "" {
|
||||
continue
|
||||
}
|
||||
slot.Name = name
|
||||
slot.AcceptedMediaTypes = append([]string(nil), slot.AcceptedMediaTypes...)
|
||||
out[name] = slot
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func validateTemplateReferenceSlots(tmpl *template.Template, declared map[string]contracts.ReferenceSlot) error {
|
||||
if tmpl == nil || tmpl.Tree == nil || tmpl.Tree.Root == nil {
|
||||
return nil
|
||||
}
|
||||
return validateReferenceNodes(tmpl.Tree.Root, declared)
|
||||
}
|
||||
|
||||
func validateReferenceNodes(node parse.Node, declared map[string]contracts.ReferenceSlot) error {
|
||||
if node == nil || reflect.ValueOf(node).IsNil() {
|
||||
return nil
|
||||
}
|
||||
switch typed := node.(type) {
|
||||
case *parse.ListNode:
|
||||
for _, child := range typed.Nodes {
|
||||
if err := validateReferenceNodes(child, declared); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
case *parse.ActionNode:
|
||||
return validateReferencePipeline(typed.Pipe, declared)
|
||||
case *parse.IfNode:
|
||||
if err := validateReferencePipeline(typed.Pipe, declared); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateReferenceNodes(typed.List, declared); err != nil {
|
||||
return err
|
||||
}
|
||||
return validateReferenceNodes(typed.ElseList, declared)
|
||||
case *parse.RangeNode:
|
||||
if err := validateReferencePipeline(typed.Pipe, declared); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateReferenceNodes(typed.List, declared); err != nil {
|
||||
return err
|
||||
}
|
||||
return validateReferenceNodes(typed.ElseList, declared)
|
||||
case *parse.WithNode:
|
||||
if err := validateReferencePipeline(typed.Pipe, declared); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateReferenceNodes(typed.List, declared); err != nil {
|
||||
return err
|
||||
}
|
||||
return validateReferenceNodes(typed.ElseList, declared)
|
||||
case *parse.TemplateNode:
|
||||
return nil
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateReferencePipeline(pipe *parse.PipeNode, declared map[string]contracts.ReferenceSlot) error {
|
||||
if pipe == nil {
|
||||
return nil
|
||||
}
|
||||
for _, cmd := range pipe.Cmds {
|
||||
if err := validateReferenceCommand(cmd, declared); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateReferenceCommand(cmd *parse.CommandNode, declared map[string]contracts.ReferenceSlot) error {
|
||||
if cmd == nil || len(cmd.Args) == 0 {
|
||||
return nil
|
||||
}
|
||||
for _, arg := range cmd.Args[1:] {
|
||||
if nested, ok := arg.(*parse.PipeNode); ok {
|
||||
if err := validateReferencePipeline(nested, declared); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
identifier, ok := cmd.Args[0].(*parse.IdentifierNode)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
if identifier.Ident != "reference" && identifier.Ident != "hasreference" {
|
||||
return nil
|
||||
}
|
||||
if len(cmd.Args) != 2 {
|
||||
return fmt.Errorf("%s requires one string slot name", identifier.Ident)
|
||||
}
|
||||
slotArg, ok := cmd.Args[1].(*parse.StringNode)
|
||||
if !ok {
|
||||
return fmt.Errorf("%s requires a string literal slot name", identifier.Ident)
|
||||
}
|
||||
slotName := strings.TrimSpace(slotArg.Text)
|
||||
if slotName == "" {
|
||||
return fmt.Errorf("%s slot name must not be empty", identifier.Ident)
|
||||
}
|
||||
if _, ok := declared[slotName]; !ok {
|
||||
return fmt.Errorf("%s slot %q is not declared", identifier.Ident, slotName)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func readPromptAsset(fsys fs.FS, assetPath string) (string, error) {
|
||||
if strings.TrimSpace(assetPath) == "" {
|
||||
return "", fmt.Errorf("prompt asset path must not be empty")
|
||||
|
||||
Reference in New Issue
Block a user