Files
audita/internal/cli/process_flags_test.go

434 lines
16 KiB
Go

package cli
import (
"io"
"reflect"
"strings"
"testing"
"gitea.maximumdirect.net/eric/audita/internal/core/config"
)
func TestProcessCLIOverridesMapsEveryConfigMutatingFlag(t *testing.T) {
tests := []struct {
name string
flagName string
value string
wantExplicitModules bool
assertOverrideFields func(t *testing.T, overrides config.CLIOverrides)
}{
{
name: "modules",
flagName: "modules",
value: "grammar,glossary",
wantExplicitModules: true,
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertStringOverride(t, "ModulesCSV", overrides.ModulesCSV, "grammar,glossary")
},
},
{
name: "output schema",
flagName: "output-schema",
value: "audita-v1",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertStringOverride(t, "OutputSchema", overrides.OutputSchema, "audita-v1")
},
},
{
name: "primary api key",
flagName: "llm-api-key",
value: "primary-key",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertStringOverride(t, "PrimaryLLMAPIKey", overrides.PrimaryLLMAPIKey, "primary-key")
},
},
{
name: "validation api key",
flagName: "validation-llm-api-key",
value: "validation-key",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertStringOverride(t, "ValidationLLMAPIKey", overrides.ValidationLLMAPIKey, "validation-key")
},
},
{
name: "primary model",
flagName: "model",
value: "primary-model",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertStringOverride(t, "PrimaryModel", overrides.PrimaryModel, "primary-model")
},
},
{
name: "validation model",
flagName: "validation-model",
value: "validation-model",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertStringOverride(t, "ValidationModel", overrides.ValidationModel, "validation-model")
},
},
{
name: "primary base url",
flagName: "base-url",
value: "https://primary.example.test",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertStringOverride(t, "PrimaryBaseURL", overrides.PrimaryBaseURL, "https://primary.example.test")
},
},
{
name: "validation base url",
flagName: "validation-base-url",
value: "https://validation.example.test",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertStringOverride(t, "ValidationBaseURL", overrides.ValidationBaseURL, "https://validation.example.test")
},
},
{
name: "primary timeout",
flagName: "llm-timeout-seconds",
value: "101",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertIntOverride(t, "PrimaryLLMTimeoutSeconds", overrides.PrimaryLLMTimeoutSeconds, 101)
},
},
{
name: "total concurrency",
flagName: "total-llm-concurrency",
value: "5",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertIntOverride(t, "TotalLLMConcurrency", overrides.TotalLLMConcurrency, 5)
},
},
{
name: "proposal concurrency",
flagName: "proposal-llm-concurrency",
value: "3",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertIntOverride(t, "ProposalLLMConcurrency", overrides.ProposalLLMConcurrency, 3)
},
},
{
name: "legacy concurrency alias",
flagName: "llm-concurrency",
value: "4",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertIntOverride(t, "PrimaryLLMConcurrency", overrides.PrimaryLLMConcurrency, 4)
},
},
{
name: "validation timeout",
flagName: "validation-llm-timeout-seconds",
value: "202",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertIntOverride(t, "ValidationLLMTimeoutSeconds", overrides.ValidationLLMTimeoutSeconds, 202)
},
},
{
name: "max retries",
flagName: "max-retries",
value: "6",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertIntOverride(t, "MaxRetries", overrides.MaxRetries, 6)
},
},
{
name: "validation max retries",
flagName: "validation-max-retries",
value: "7",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertIntOverride(t, "ValidationMaxRetries", overrides.ValidationMaxRetries, 7)
},
},
{
name: "validation concurrency",
flagName: "validation-llm-concurrency",
value: "8",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertIntOverride(t, "ValidationLLMConcurrency", overrides.ValidationLLMConcurrency, 8)
},
},
{
name: "validation max prompt tokens",
flagName: "validation-max-prompt-tokens",
value: "4096",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertIntOverride(t, "ValidationMaxPromptTokens", overrides.ValidationMaxPromptTokens, 4096)
},
},
{
name: "max section tokens",
flagName: "max-section-tokens",
value: "9000",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertIntOverride(t, "MaxSectionTokens", overrides.MaxSectionTokens, 9000)
},
},
{
name: "min section tokens",
flagName: "min-section-tokens",
value: "1000",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertIntOverride(t, "MinSectionTokens", overrides.MinSectionTokens, 1000)
},
},
{
name: "target sections",
flagName: "target-sections",
value: "12",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertIntOverride(t, "TargetSections", overrides.TargetSections, 12)
},
},
{
name: "glossary threshold",
flagName: "glossary-confidence-threshold",
value: "0.91",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertFloatOverride(t, "GlossaryConfidenceThreshold", overrides.GlossaryConfidenceThreshold, 0.91)
},
},
{
name: "grammar threshold",
flagName: "grammar-confidence-threshold",
value: "0.92",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertFloatOverride(t, "GrammarConfidenceThreshold", overrides.GrammarConfidenceThreshold, 0.92)
},
},
{
name: "homophones threshold",
flagName: "homophones-confidence-threshold",
value: "0.93",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertFloatOverride(t, "HomophonesConfidenceThreshold", overrides.HomophonesConfidenceThreshold, 0.93)
},
},
{
name: "spoken word threshold",
flagName: "spoken-word-confidence-threshold",
value: "0.94",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertFloatOverride(t, "SpokenWordConfidenceThreshold", overrides.SpokenWordConfidenceThreshold, 0.94)
},
},
{
name: "normalize max segment gap",
flagName: "normalize-max-segment-gap",
value: "1.2",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertFloatOverride(t, "NormalizeMaxSegmentGap", overrides.NormalizeMaxSegmentGap, 1.2)
},
},
{
name: "normalize ellipsis gap",
flagName: "normalize-ellipsis-gap",
value: "2.3",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertFloatOverride(t, "NormalizeEllipsisGap", overrides.NormalizeEllipsisGap, 2.3)
},
},
{
name: "normalize max segment duration",
flagName: "normalize-max-segment-duration",
value: "45.6",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertFloatOverride(t, "NormalizeMaxSegmentDuration", overrides.NormalizeMaxSegmentDuration, 45.6)
},
},
{
name: "normalize max segment tokens",
flagName: "normalize-max-segment-tokens",
value: "321",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertIntOverride(t, "NormalizeMaxSegmentTokens", overrides.NormalizeMaxSegmentTokens, 321)
},
},
{
name: "transcript description",
flagName: "transcript-description",
value: "podcast episode",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertStringOverride(t, "TranscriptDescription", overrides.TranscriptDescription, "podcast episode")
},
},
{
name: "work dir",
flagName: "work-dir",
value: "/tmp/custom-audita",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertStringOverride(t, "WorkDir", overrides.WorkDir, "/tmp/custom-audita")
},
},
{
name: "work dir retention",
flagName: "work-dir-retention",
value: "always",
assertOverrideFields: func(t *testing.T, overrides config.CLIOverrides) {
assertStringOverride(t, "WorkDirRetention", overrides.WorkDirRetention, "always")
},
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
fs, flags := newProcessFlagSet(config.Default(), io.Discard)
if err := fs.Parse([]string{"--" + tc.flagName, tc.value}); err != nil {
t.Fatalf("parse flag: %v", err)
}
overrides, explicitModules := processCLIOverrides(fs, flags)
if explicitModules != tc.wantExplicitModules {
t.Fatalf("explicitModules=%v, want %v", explicitModules, tc.wantExplicitModules)
}
tc.assertOverrideFields(t, overrides)
})
}
}
func TestProcessCLIOverridesIgnoresNonConfigFlags(t *testing.T) {
fs, flags := newProcessFlagSet(config.Default(), io.Discard)
if err := fs.Parse([]string{
"--config", "/tmp/config.yml",
"--glossary", "/tmp/glossary.yml",
"--output", "/tmp/output.json",
"--report-json", "/tmp/report.json",
}); err != nil {
t.Fatalf("parse flags: %v", err)
}
overrides, explicitModules := processCLIOverrides(fs, flags)
if explicitModules {
t.Fatal("non-config flags should not mark modules explicit")
}
assertNoCLIOverrides(t, overrides)
}
func TestNewProcessFlagSetDefaultsReflectEffectiveConfig(t *testing.T) {
cfg := config.Default()
cfg.Modules = []string{"grammar", "glossary"}
cfg.OutputSchema = "audita-v1"
cfg.PrimaryLLM.APIKey = "primary-key"
cfg.ValidationLLM.APIKey = "validation-key"
cfg.PrimaryLLM.Model = "primary-model"
cfg.ValidationLLM.Model = "validation-model"
cfg.PrimaryLLM.BaseURL = "https://primary.example.test"
cfg.ValidationLLM.BaseURL = "https://validation.example.test"
cfg.PrimaryLLM.TimeoutSeconds = 101
cfg.TotalLLMConcurrency = 5
cfg.ProposalLLMConcurrency = 3
cfg.PrimaryLLM.MaxRetries = 6
cfg.ValidationMaxPromptTokens = 4096
cfg.MaxSectionTokens = 9000
cfg.MinSectionTokens = 1000
cfg.Thresholds.Glossary = 0.91
cfg.Thresholds.Grammar = 0.92
cfg.Thresholds.Homophones = 0.93
cfg.Thresholds.SpokenWord = 0.94
cfg.Normalization.MaxSegmentGap = 1.2
cfg.Normalization.EllipsisGap = 2.3
cfg.Normalization.MaxSegmentDuration = 45.6
cfg.Normalization.MaxSegmentTokens = 321
cfg.TranscriptDescription = "podcast episode"
cfg.WorkDir = "/tmp/custom-audita"
cfg.WorkDirRetention = config.WorkDirRetentionAlways
validationTimeout := 202
validationRetries := 7
validationConcurrency := 8
targetSections := 12
cfg.ValidationLLM.TimeoutSeconds = &validationTimeout
cfg.ValidationLLM.MaxRetries = &validationRetries
cfg.ValidationLLMConcurrency = &validationConcurrency
cfg.TargetSections = &targetSections
_, flags := newProcessFlagSet(cfg, io.Discard)
assertStringOverride(t, "modules default", flags.modules, "grammar,glossary")
assertStringOverride(t, "output schema default", flags.outputSchema, "audita-v1")
assertStringOverride(t, "primary api key default", flags.llmAPIKey, "primary-key")
assertStringOverride(t, "validation api key default", flags.validationLLMAPIKey, "validation-key")
assertStringOverride(t, "primary model default", flags.model, "primary-model")
assertStringOverride(t, "validation model default", flags.validationModel, "validation-model")
assertStringOverride(t, "primary base url default", flags.baseURL, "https://primary.example.test")
assertStringOverride(t, "validation base url default", flags.validationBaseURL, "https://validation.example.test")
assertIntOverride(t, "primary timeout default", flags.llmTimeoutSeconds, 101)
assertIntOverride(t, "total concurrency default", flags.totalLLMConcurrency, 5)
assertIntOverride(t, "proposal concurrency default", flags.proposalLLMConcurrency, 3)
assertIntOverride(t, "legacy concurrency alias default", flags.llmConcurrency, 5)
assertIntOverride(t, "validation timeout default", flags.validationLLMTimeoutSeconds, validationTimeout)
assertIntOverride(t, "max retries default", flags.maxRetries, 6)
assertIntOverride(t, "validation max retries default", flags.validationMaxRetries, validationRetries)
assertIntOverride(t, "validation concurrency default", flags.validationLLMConcurrency, validationConcurrency)
assertIntOverride(t, "validation max prompt tokens default", flags.validationMaxPromptTokens, 4096)
assertIntOverride(t, "max section tokens default", flags.maxSectionTokens, 9000)
assertIntOverride(t, "min section tokens default", flags.minSectionTokens, 1000)
assertIntOverride(t, "target sections default", flags.targetSections, targetSections)
assertFloatOverride(t, "glossary threshold default", flags.glossaryConfidenceThreshold, 0.91)
assertFloatOverride(t, "grammar threshold default", flags.grammarConfidenceThreshold, 0.92)
assertFloatOverride(t, "homophones threshold default", flags.homophonesConfidenceThreshold, 0.93)
assertFloatOverride(t, "spoken word threshold default", flags.spokenWordConfidenceThreshold, 0.94)
assertFloatOverride(t, "normalize max segment gap default", flags.normalizeMaxSegmentGap, 1.2)
assertFloatOverride(t, "normalize ellipsis gap default", flags.normalizeEllipsisGap, 2.3)
assertFloatOverride(t, "normalize max segment duration default", flags.normalizeMaxSegmentDuration, 45.6)
assertIntOverride(t, "normalize max segment tokens default", flags.normalizeMaxSegmentTokens, 321)
assertStringOverride(t, "transcript description default", flags.transcriptDescription, "podcast episode")
assertStringOverride(t, "work dir default", flags.workDir, "/tmp/custom-audita")
assertStringOverride(t, "work dir retention default", flags.workDirRetention, "always")
}
func TestNewProcessFlagSetUsesFallbackDefaultsForUnsetOptionalConfig(t *testing.T) {
cfg := config.Default()
_, flags := newProcessFlagSet(cfg, io.Discard)
assertIntOverride(t, "validation timeout fallback", flags.validationLLMTimeoutSeconds, cfg.PrimaryLLM.TimeoutSeconds)
assertIntOverride(t, "validation retries fallback", flags.validationMaxRetries, cfg.PrimaryLLM.MaxRetries)
assertIntOverride(t, "validation concurrency fallback", flags.validationLLMConcurrency, cfg.TotalLLMConcurrency)
assertIntOverride(t, "target sections fallback", flags.targetSections, 0)
}
func assertStringOverride(t *testing.T, name string, got *string, want string) {
t.Helper()
if got == nil || *got != want {
t.Fatalf("%s=%v, want %q", name, pointerValue(got), want)
}
}
func assertIntOverride(t *testing.T, name string, got *int, want int) {
t.Helper()
if got == nil || *got != want {
t.Fatalf("%s=%v, want %d", name, pointerValue(got), want)
}
}
func assertFloatOverride(t *testing.T, name string, got *float64, want float64) {
t.Helper()
if got == nil || *got != want {
t.Fatalf("%s=%v, want %v", name, pointerValue(got), want)
}
}
func assertNoCLIOverrides(t *testing.T, overrides config.CLIOverrides) {
t.Helper()
value := reflect.ValueOf(overrides)
typ := value.Type()
for i := 0; i < value.NumField(); i++ {
field := value.Field(i)
if field.Kind() != reflect.Ptr {
t.Fatalf("unexpected non-pointer CLIOverrides field %s", typ.Field(i).Name)
}
if !field.IsNil() {
t.Fatalf("expected no CLI overrides, field %s was set", typ.Field(i).Name)
}
}
}
func pointerValue[T any](ptr *T) any {
if ptr == nil {
return "<nil>"
}
if stringer, ok := any(*ptr).(interface{ String() string }); ok {
return strings.TrimSpace(stringer.String())
}
return *ptr
}