Expose run-wide reasoning effort controls

This commit is contained in:
2026-07-30 02:11:35 +00:00
parent f603f7ac64
commit f8333f2c15
9 changed files with 269 additions and 47 deletions

View File

@@ -12,6 +12,7 @@ import (
"gitea.maximumdirect.net/eric/notarius/internal/core/artifacts"
"gitea.maximumdirect.net/eric/notarius/internal/core/config"
"gitea.maximumdirect.net/eric/notarius/internal/framework/checkpoint"
"gitea.maximumdirect.net/eric/notarius/internal/framework/contracts"
"gitea.maximumdirect.net/eric/notarius/internal/framework/pipeline"
)
@@ -319,6 +320,129 @@ func TestRunLLMProfileOverrideAndValidationUseInjectedBoundaries(t *testing.T) {
})
}
func TestRunReasoningEffortOverrideReachesFactory(t *testing.T) {
tests := []struct {
name string
flags []string
wantValue string
wantSet bool
}{
{name: "inherit"},
{name: "replace", flags: []string{"--reasoning-effort", " focused "}, wantValue: "focused", wantSet: true},
{name: "clear", flags: []string{"--clear-reasoning-effort"}, wantSet: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
roots := newStateTestRoots(t)
opts := newStateTestHarness().options()
var got []LLMRuntimeOverrides
opts.LLMClientFactory = func(_ context.Context, _ config.Config, _ string, overrides LLMRuntimeOverrides) (contracts.StructuredLLMClient, []artifacts.LLMProfileManifest, error) {
got = append(got, overrides)
return nil, nil, nil
}
args := append([]string{"run", "sample", "--config", roots.config, "--input", roots.input, "--chunk_cache", "bypass"}, tt.flags...)
var stdout, stderr bytes.Buffer
if code := RunWithOptions(args, &stdout, &stderr, opts); code != 0 || stderr.Len() != 0 {
t.Fatalf("code=%d stdout=%q stderr=%q", code, stdout.String(), stderr.String())
}
if len(got) != 1 {
t.Fatalf("factory overrides = %#v, want one call", got)
}
if !tt.wantSet {
if got[0].ReasoningEffort != nil {
t.Fatalf("reasoning effort = %q, want inherit", *got[0].ReasoningEffort)
}
return
}
if got[0].ReasoningEffort == nil || *got[0].ReasoningEffort != tt.wantValue {
t.Fatalf("reasoning effort = %#v, want %q", got[0].ReasoningEffort, tt.wantValue)
}
})
}
}
func TestRunReasoningEffortOverrideRejectsInvalidSyntax(t *testing.T) {
tests := []struct {
name string
flags []string
wantError string
}{
{
name: "mutually exclusive controls",
flags: []string{"--reasoning-effort", "focused", "--clear-reasoning-effort"},
wantError: "cannot be combined",
},
{
name: "empty replacement",
flags: []string{"--reasoning-effort", " "},
wantError: "must not be empty",
},
{
name: "duplicate replacement",
flags: []string{"--reasoning-effort", "low", "--reasoning-effort", "high"},
wantError: "may be specified only once",
},
{
name: "missing replacement",
flags: []string{"--reasoning-effort"},
wantError: "flag needs an argument",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
roots := newStateTestRoots(t)
args := append([]string{"run", "sample", "--config", roots.config, "--input", roots.input}, tt.flags...)
var stdout, stderr bytes.Buffer
code := RunWithOptions(args, &stdout, &stderr, newStateTestHarness().options())
if code != 2 || stdout.Len() != 0 || !strings.Contains(stderr.String(), tt.wantError) {
t.Fatalf("code=%d stdout=%q stderr=%q", code, stdout.String(), stderr.String())
}
assertNoRunState(t, roots)
})
}
}
func TestReasoningEffortOverrideSeparatesCheckpointIdentities(t *testing.T) {
replacement := " focused "
cleared := ""
states := []struct {
name string
overrides LLMRuntimeOverrides
wantValue string
wantSet bool
}{
{name: "inherit"},
{name: "replace", overrides: LLMRuntimeOverrides{ReasoningEffort: &replacement}, wantValue: "focused", wantSet: true},
{name: "clear", overrides: LLMRuntimeOverrides{ReasoningEffort: &cleared}, wantValue: "<cleared>", wantSet: true},
}
digests := make(map[string]string, len(states))
for _, state := range states {
fingerprints := runtimeOverrideFingerprints("", "", state.overrides)
var value string
var found bool
for _, fingerprint := range fingerprints {
if fingerprint.Name == "reasoning_effort_override" {
value, found = fingerprint.Value, true
}
}
if found != state.wantSet || (found && value != state.wantValue) {
t.Fatalf("%s fingerprint found=%t value=%q, want found=%t value=%q", state.name, found, value, state.wantSet, state.wantValue)
}
identity, err := checkpoint.NewIdentity(checkpoint.IdentityInput{
Pipeline: pipeline.ResolvedPipeline{ID: "sample", Digest: "sha256:pipeline", Input: pipeline.Binding("test/input")},
RawInputDigest: "sha256:input",
RuntimeOverrides: fingerprints,
})
if err != nil {
t.Fatal(err)
}
digests[state.name] = identity.Digest
}
if digests["inherit"] == digests["replace"] || digests["inherit"] == digests["clear"] || digests["replace"] == digests["clear"] {
t.Fatalf("checkpoint identity digests are not distinct: %#v", digests)
}
}
func TestEffectiveLLMProfileIDsAreSortedDeduplicatedAndLLMOnly(t *testing.T) {
resolved := pipeline.ResolvedPipeline{
Input: pipeline.ModuleBinding{LLMProfile: "input-profile"},