434 lines
16 KiB
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
|
|
}
|