Files
narratio/internal/config/load.go
Eric Rakestraw 33f7ae8f2e
All checks were successful
ci/woodpecker/tag/release Pipeline was successful
Simplify downstream tool configuration
2026-05-16 23:30:40 +00:00

327 lines
7.7 KiB
Go

package config
import (
"fmt"
"io"
"os"
"path/filepath"
"regexp"
"strings"
"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) {
return LoadSessionWithOptions(path, SessionLoadOptions{})
}
// SessionLoadOptions configures session template rendering behavior.
type SessionLoadOptions struct {
SessionID string
}
// LoadSessionWithOptions loads session configuration from a YAML file with
// strict field checking after template rendering.
func LoadSessionWithOptions(path string, opts SessionLoadOptions) (*SessionConfig, error) {
sessionBytes, err := os.ReadFile(path)
if err != nil {
return nil, fmt.Errorf("load session config: session file %q: open: %w", path, err)
}
rendered, err := renderSessionTemplate(string(sessionBytes), opts)
if err != nil {
return nil, fmt.Errorf("load session config: %w", err)
}
var cfg SessionConfig
if err := decodeStrictYAMLFromReader("session", path, strings.NewReader(rendered), &cfg); err != nil {
return nil, fmt.Errorf("load session config: %w", err)
}
if strings.TrimSpace(opts.SessionID) != "" && strings.TrimSpace(cfg.SessionID) != "" && strings.TrimSpace(cfg.SessionID) != strings.TrimSpace(opts.SessionID) {
return nil, fmt.Errorf(
"load session config: session file %q: session_id mismatch: --session-id %q does not match rendered session_id %q",
path,
strings.TrimSpace(opts.SessionID),
strings.TrimSpace(cfg.SessionID),
)
}
return &cfg, nil
}
// Load loads and resolves combined pipeline and session configuration.
func Load(pipelinePath, sessionPath string) (*Config, error) {
return LoadWithSessionOptions(pipelinePath, sessionPath, SessionLoadOptions{})
}
// LoadWithSessionOptions loads and resolves combined pipeline and session
// configuration with session template options.
func LoadWithSessionOptions(pipelinePath, sessionPath string, sessionOpts SessionLoadOptions) (*Config, error) {
pipelineCfg, err := LoadPipeline(pipelinePath)
if err != nil {
return nil, err
}
sessionCfg, err := LoadSessionWithOptions(sessionPath, sessionOpts)
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()
return decodeStrictYAMLFromReader(kind, path, f, out)
}
func decodeStrictYAMLFromReader(kind, path string, r io.Reader, out any) error {
dec := yaml.NewDecoder(r)
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
}
var sessionTemplatePattern = regexp.MustCompile(`\{\{\s*([a-zA-Z_][a-zA-Z0-9_]*)\s*\}\}`)
func renderSessionTemplate(content string, opts SessionLoadOptions) (string, error) {
sessionID := strings.TrimSpace(opts.SessionID)
rendered := content
if sessionID != "" {
rendered = strings.ReplaceAll(rendered, "{{session_id}}", sessionID)
rendered = strings.ReplaceAll(rendered, "{{ session_id }}", sessionID)
}
unresolved := sessionTemplatePattern.FindAllStringSubmatch(rendered, -1)
if len(unresolved) > 0 {
vars := make([]string, 0, len(unresolved))
for _, m := range unresolved {
if len(m) > 1 {
vars = append(vars, m[1])
}
}
if len(vars) > 0 {
return "", fmt.Errorf(
"session file template rendering failed: unresolved template variable(s): %s; pass --session-id when using {{ session_id }}",
strings.Join(vars, ", "),
)
}
return "", fmt.Errorf("session file template rendering failed: unresolved template placeholders remain")
}
return rendered, 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
}
applyStorageDefaults(&cfg.Storage)
applySpoolDefaults(&cfg.Spool)
applyArchiveDefaults(&cfg.Archive)
applyWhisperXDefaults(&cfg.WhisperX)
applySeriatimDefaults(&cfg.Seriatim)
applyAuditaDefaults(&cfg.Audita)
if cfg.Normalize == nil {
cfg.Normalize = &NormalizeConfig{}
}
applyNormalizeDefaults(cfg.Normalize)
applyTrimDefaults(cfg.Trim)
applyScriptoriumDefaults(cfg.Scriptorium)
}
func applyStorageDefaults(cfg *StorageConfig) {
if cfg == nil {
return
}
if cfg.S3 == nil {
cfg.S3 = &StorageS3Config{}
}
if cfg.S3.RootPrefix == "" {
cfg.S3.RootPrefix = "dnd"
}
}
func applySpoolDefaults(cfg *SpoolConfig) {
if cfg == nil {
return
}
if cfg.Root == "" {
cfg.Root = "/var/spool/narratio"
}
}
func applyArchiveDefaults(cfg **ArchiveConfig) {
if cfg == nil {
return
}
if *cfg == nil {
*cfg = &ArchiveConfig{}
}
if (*cfg).Enabled == nil {
(*cfg).Enabled = boolPtr(true)
}
if (*cfg).UploadRun == nil {
(*cfg).UploadRun = boolPtr(true)
}
if len((*cfg).PromoteArtifacts) == 0 {
(*cfg).PromoteArtifacts = []ArchivePromotionRule{
{From: "transcripts/trimmed.json", To: "transcripts/trimmed.json", Required: boolPtr(true)},
{From: "artifacts/session_recap.md", To: "artifacts/session_recap.md", Required: boolPtr(true)},
}
}
for i := range (*cfg).PromoteArtifacts {
if (*cfg).PromoteArtifacts[i].Required == nil {
(*cfg).PromoteArtifacts[i].Required = boolPtr(true)
}
}
}
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.Binary == "" {
cfg.Binary = "seriatim"
}
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.Binary == "" {
cfg.Binary = "audita"
}
if cfg.Timeout == "" {
cfg.Timeout = "3h"
}
if cfg.Report == nil {
cfg.Report = boolPtr(true)
}
}
func applyScriptoriumDefaults(cfg *ScriptoriumConfig) {
if cfg == nil {
return
}
if cfg.Binary == "" {
cfg.Binary = "scriptorium"
}
if cfg.Timeout == "" {
cfg.Timeout = "10m"
}
}
func applyTrimDefaults(cfg *TrimConfig) {
if cfg == nil {
return
}
if cfg.Bounds.Timeout == "" {
cfg.Bounds.Timeout = "10m"
}
if cfg.Seriatim.Report == nil {
cfg.Seriatim.Report = boolPtr(false)
}
}
func applyNormalizeDefaults(cfg *NormalizeConfig) {
if cfg == nil {
return
}
if cfg.OutputSchema == "" {
cfg.OutputSchema = defaultNormalizeOutputSchema
}
if cfg.OutputPath == "" && !cfg.outputPathWasSet() {
cfg.OutputPath = defaultNormalizeOutputPath
}
if cfg.Report == nil {
cfg.Report = boolPtr(true)
}
}
func float64Ptr(v float64) *float64 {
p := v
return &p
}
func boolPtr(v bool) *bool {
p := v
return &p
}