summaryrefslogtreecommitdiff
path: root/internal/config/config.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/config/config.go')
-rw-r--r--internal/config/config.go22
1 files changed, 18 insertions, 4 deletions
diff --git a/internal/config/config.go b/internal/config/config.go
index 9267ff0..2726c46 100644
--- a/internal/config/config.go
+++ b/internal/config/config.go
@@ -5,6 +5,7 @@ package config
import (
"errors"
"fmt"
+ "io/fs"
"os"
"path/filepath"
"strings"
@@ -13,6 +14,7 @@ import (
"github.com/spf13/viper"
"codeberg.org/snonux/comicforge/internal/provider"
+ "codeberg.org/snonux/comicforge/prompts"
)
const (
@@ -228,16 +230,19 @@ func (c *Config) PromptPath(name string) string {
func (c *Config) LoadPromptTemplate(name string) (*template.Template, error) {
path := c.PromptPath(name)
content, err := os.ReadFile(path)
- if err != nil {
+ if err == nil {
+ return parsePromptTemplate(path, content)
+ }
+ if !errors.Is(err, fs.ErrNotExist) {
return nil, fmt.Errorf("read prompt %q: %w", path, err)
}
- tmpl, err := template.New(filepath.Base(path)).Option("missingkey=error").Parse(string(content))
+ content, err = prompts.Read(name)
if err != nil {
- return nil, fmt.Errorf("parse prompt %q: %w", path, err)
+ return nil, fmt.Errorf("read embedded prompt %q: %w", name, err)
}
- return tmpl, nil
+ return parsePromptTemplate(filepath.Base(name), content)
}
// RenderPrompt executes a prompt template with the provided data.
@@ -255,6 +260,15 @@ func (c *Config) RenderPrompt(name string, data any) (string, error) {
return builder.String(), nil
}
+func parsePromptTemplate(name string, content []byte) (*template.Template, error) {
+ tmpl, err := template.New(filepath.Base(name)).Option("missingkey=error").Parse(string(content))
+ if err != nil {
+ return nil, fmt.Errorf("parse prompt %q: %w", name, err)
+ }
+
+ return tmpl, nil
+}
+
func (c *Config) normalize() {
c.Provider.Text = provider.NormalizeName(c.Provider.Text)
c.Provider.Image = provider.NormalizeName(c.Provider.Image)