Files
narratio/internal/config/load.go

187 lines
3.7 KiB
Go

package config
import (
"fmt"
"io"
"os"
"path/filepath"
"gopkg.in/yaml.v3"
)
// LoadPipeline loads pipeline configuration from a YAML file with strict field checking.
func LoadPipeline(path string) (*PipelineConfig, error) {
var cfg PipelineConfig
if err := decodeStrictYAML("pipeline", path, &cfg); err != nil {
return nil, fmt.Errorf("load pipeline config: %w", err)
}
applyPipelineDefaults(&cfg)
return &cfg, nil
}
// LoadSession loads session configuration from a YAML file with strict field checking.
func LoadSession(path string) (*SessionConfig, error) {
var cfg SessionConfig
if err := decodeStrictYAML("session", path, &cfg); err != nil {
return nil, fmt.Errorf("load session config: %w", err)
}
return &cfg, nil
}
// Load loads and resolves combined pipeline and session configuration.
func Load(pipelinePath, sessionPath string) (*Config, error) {
pipelineCfg, err := LoadPipeline(pipelinePath)
if err != nil {
return nil, err
}
sessionCfg, err := LoadSession(sessionPath)
if err != nil {
return nil, err
}
return &Config{
Pipeline: pipelineCfg,
Session: sessionCfg,
PipelinePath: pipelinePath,
SessionPath: sessionPath,
}, nil
}
func decodeStrictYAML(kind, path string, out any) error {
f, err := os.Open(path)
if err != nil {
return fmt.Errorf("%s file %q: open: %w", kind, path, err)
}
defer f.Close()
dec := yaml.NewDecoder(f)
dec.KnownFields(true)
if err := dec.Decode(out); err != nil {
return fmt.Errorf("%s file %q: strict decode failed: %w", kind, path, err)
}
var extra any
if err := dec.Decode(&extra); err != nil && err != io.EOF {
return fmt.Errorf("%s file %q: trailing content decode failed: %w", kind, path, err)
}
return nil
}
func shortName(path, fallback string) string {
base := filepath.Base(path)
if base == "." || base == string(filepath.Separator) {
return fallback
}
return base
}
func applyPipelineDefaults(cfg *PipelineConfig) {
if cfg == nil {
return
}
applyWhisperXDefaults(&cfg.WhisperX)
applySeriatimDefaults(&cfg.Seriatim)
applyAuditaDefaults(&cfg.Audita)
applyScriptoriumDefaults(cfg.Scriptorium)
}
func applyWhisperXDefaults(cfg *WhisperXConfig) {
if cfg == nil {
return
}
if cfg.Language == "" {
cfg.Language = "en"
}
if cfg.Timeout == "" {
cfg.Timeout = "30m"
}
if cfg.RetryDelay == "" {
cfg.RetryDelay = "2s"
}
if cfg.Concurrency == nil {
cfg.Concurrency = intPtr(2)
}
if cfg.Retries == nil {
cfg.Retries = intPtr(3)
}
}
func intPtr(v int) *int {
p := v
return &p
}
func applySeriatimDefaults(cfg *SeriatimConfig) {
if cfg == nil {
return
}
if cfg.Timeout == "" {
cfg.Timeout = "10m"
}
if cfg.OutputSchema == "" {
cfg.OutputSchema = "seriatim-intermediate"
}
if cfg.CoalesceGap == nil {
cfg.CoalesceGap = float64Ptr(3.0)
}
if cfg.Report == nil {
cfg.Report = boolPtr(true)
}
}
func applyAuditaDefaults(cfg *AuditaConfig) {
if cfg == nil {
return
}
if cfg.Timeout == "" {
cfg.Timeout = "3h"
}
if cfg.Modules == nil {
cfg.Modules = []string{
"glossary",
"homophones",
"glossary",
"spoken_word",
"grammar",
"homophones",
"glossary",
}
}
if cfg.BaseURL == "" {
cfg.BaseURL = "https://openrouter.ai/api/v1"
}
if cfg.Model == "" {
cfg.Model = "openrouter/google/gemma-4-31b-it"
}
if cfg.LLMConcurrency == nil {
cfg.LLMConcurrency = intPtr(1)
}
if cfg.ValidationLLMConcurrency == nil {
cfg.ValidationLLMConcurrency = intPtr(1)
}
if cfg.Report == nil {
cfg.Report = boolPtr(true)
}
}
func applyScriptoriumDefaults(cfg *ScriptoriumConfig) {
if cfg == nil {
return
}
if cfg.Timeout == "" {
cfg.Timeout = "10m"
}
}
func float64Ptr(v float64) *float64 {
p := v
return &p
}
func boolPtr(v bool) *bool {
p := v
return &p
}