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 }