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.go52
1 files changed, 49 insertions, 3 deletions
diff --git a/internal/config/config.go b/internal/config/config.go
index ba8a56e..5b4ab18 100644
--- a/internal/config/config.go
+++ b/internal/config/config.go
@@ -13,7 +13,10 @@ import (
"github.com/spf13/viper"
+ "codeberg.org/snonux/comicforge/internal/image"
"codeberg.org/snonux/comicforge/internal/provider"
+ "codeberg.org/snonux/comicforge/internal/text"
+ "codeberg.org/snonux/comicforge/internal/tts"
"codeberg.org/snonux/comicforge/prompts"
)
@@ -40,6 +43,9 @@ var (
_ provider.TextConfig = (*Config)(nil)
_ provider.ImageConfig = (*Config)(nil)
_ provider.TTSConfig = (*Config)(nil)
+ _ text.Config = (*Config)(nil)
+ _ image.Config = (*Config)(nil)
+ _ tts.Config = (*Config)(nil)
)
// ProviderConfig stores the selected provider name for each capability.
@@ -208,6 +214,46 @@ func (c *Config) ImageProviderName() string {
return provider.NormalizeName(c.Provider.Image)
}
+// GoogleAPIKey returns the configured Google API key.
+func (c *Config) GoogleAPIKey() string {
+ if c == nil {
+ return ""
+ }
+ return c.API.GoogleAPIKey
+}
+
+// TextModel returns the configured text model name.
+func (c *Config) TextModel() string {
+ if c == nil {
+ return ""
+ }
+ return c.Models.Text
+}
+
+// ImageModel returns the configured image model name.
+func (c *Config) ImageModel() string {
+ if c == nil {
+ return ""
+ }
+ return c.Models.Image
+}
+
+// ImageTextModel returns the configured image-text model name.
+func (c *Config) ImageTextModel() string {
+ if c == nil {
+ return ""
+ }
+ return c.Models.ImageText
+}
+
+// TTSModel returns the configured text-to-speech model name.
+func (c *Config) TTSModel() string {
+ if c == nil {
+ return ""
+ }
+ return c.Models.TTS
+}
+
// TTSProviderName returns the configured TTS provider name.
func (c *Config) TTSProviderName() string {
return provider.NormalizeName(c.Provider.TTS)
@@ -290,13 +336,13 @@ func (c *Config) normalize() {
}
func (c *Config) validate() error {
- if !provider.IsKnownName(c.Provider.Text) {
+ if !text.DefaultRegistry().Has(c.Provider.Text) {
return fmt.Errorf("unknown text provider: %s", c.Provider.Text)
}
- if !provider.IsKnownName(c.Provider.Image) {
+ if !image.DefaultRegistry().Has(c.Provider.Image) {
return fmt.Errorf("unknown image provider: %s", c.Provider.Image)
}
- if !provider.IsKnownName(c.Provider.TTS) {
+ if !tts.DefaultRegistry().Has(c.Provider.TTS) {
return fmt.Errorf("unknown TTS provider: %s", c.Provider.TTS)
}