diff --git a/internal/adapter/cli/run.go b/internal/adapter/cli/run.go index 99c643a..471f608 100644 --- a/internal/adapter/cli/run.go +++ b/internal/adapter/cli/run.go @@ -14,6 +14,7 @@ import ( httpadapter "gitea.maximumdirect.net/eric/scriptorium/internal/adapter/http" artifactadapter "gitea.maximumdirect.net/eric/scriptorium/internal/artifact" + "gitea.maximumdirect.net/eric/scriptorium/internal/defaults" "gitea.maximumdirect.net/eric/scriptorium/internal/domain" "gitea.maximumdirect.net/eric/scriptorium/internal/llm" "gitea.maximumdirect.net/eric/scriptorium/internal/profile" @@ -224,8 +225,8 @@ func parseRunArgs(args []string) (*runConfig, error) { fs.StringVar(&cfg.model, "model", "", "model name") fs.Float64Var(&cfg.temperature, "temperature", 0, "optional temperature override") fs.IntVar(&cfg.maxTokens, "max-tokens", 0, "optional max tokens override") - fs.StringVar(&cfg.schemaDir, "schema-dir", ".", "base directory for validation schemas") - fs.DurationVar(&cfg.timeout, "timeout", 10*time.Minute, "LLM request timeout") + fs.StringVar(&cfg.schemaDir, "schema-dir", defaults.SchemaDirDefault, "base directory for validation schemas") + fs.DurationVar(&cfg.timeout, "timeout", defaults.LLMRequestTimeoutDefault, "LLM request timeout") if err := fs.Parse(args); err != nil { return nil, err @@ -262,12 +263,12 @@ func parseServeArgs(args []string) (*serveConfig, error) { fs := flag.NewFlagSet("serve", flag.ContinueOnError) fs.SetOutput(io.Discard) - fs.StringVar(&cfg.addr, "addr", ":8080", "HTTP listen address") + fs.StringVar(&cfg.addr, "addr", defaults.HTTPAddrDefault, "HTTP listen address") fs.StringVar(&cfg.profileDir, "profile-dir", "", "directory containing prompt definition YAML files") - fs.StringVar(&cfg.schemaDir, "schema-dir", ".", "base directory for validation schemas") + fs.StringVar(&cfg.schemaDir, "schema-dir", defaults.SchemaDirDefault, "base directory for validation schemas") fs.StringVar(&cfg.llmBaseURL, "llm-base-url", "", "OpenAI-compatible base URL including /v1") fs.StringVar(&cfg.model, "model", "", "optional default model") - fs.DurationVar(&cfg.timeout, "timeout", 10*time.Minute, "LLM request timeout") + fs.DurationVar(&cfg.timeout, "timeout", defaults.LLMRequestTimeoutDefault, "LLM request timeout") if err := fs.Parse(args); err != nil { return nil, err diff --git a/internal/defaults/defaults.go b/internal/defaults/defaults.go new file mode 100644 index 0000000..9dc8803 --- /dev/null +++ b/internal/defaults/defaults.go @@ -0,0 +1,34 @@ +package defaults + +import ( + "time" + + "gitea.maximumdirect.net/eric/scriptorium/internal/domain" +) + +const ( + HTTPAddrDefault = ":8080" + SchemaDirDefault = "." + OutputArtifactName = "output" + ContentTypeTextPlain = "text/plain" + ContentTypeTextMarkdown = "text/markdown" + ContentTypeApplicationJSON = "application/json" + + ExecutionDefaultTemperature = 0.0 + ExecutionDefaultMaxTokens = 0 + ExecutionDefaultTopP = 1.0 + ExecutionDefaultTimeoutSeconds = 600 +) + +var ( + LLMRequestTimeoutDefault = 10 * time.Minute +) + +func ExecutionTargetDefault() domain.ExecutionTarget { + return domain.ExecutionTarget{ + Temperature: ExecutionDefaultTemperature, + MaxTokens: ExecutionDefaultMaxTokens, + TopP: ExecutionDefaultTopP, + TimeoutSeconds: ExecutionDefaultTimeoutSeconds, + } +} diff --git a/internal/llm/openai_compatible_client.go b/internal/llm/openai_compatible_client.go index e02f257..5e98ffe 100644 --- a/internal/llm/openai_compatible_client.go +++ b/internal/llm/openai_compatible_client.go @@ -7,12 +7,13 @@ import ( "errors" "fmt" "io" - "os" "net/http" "net/url" + "os" "strings" "time" + "gitea.maximumdirect.net/eric/scriptorium/internal/defaults" "gitea.maximumdirect.net/eric/scriptorium/internal/domain" ) @@ -48,7 +49,7 @@ func NewOpenAICompatibleClient(cfg OpenAICompatibleConfig) (*OpenAICompatibleCli timeout := cfg.Timeout if timeout <= 0 { - timeout = 10 * time.Minute + timeout = defaults.LLMRequestTimeoutDefault } var client *http.Client diff --git a/internal/usecase/runner.go b/internal/usecase/runner.go index 460136b..dfcc9a1 100644 --- a/internal/usecase/runner.go +++ b/internal/usecase/runner.go @@ -8,10 +8,12 @@ import ( "encoding/json" "errors" "fmt" + "os" "strings" "time" "gitea.maximumdirect.net/eric/scriptorium/internal/artifact" + "gitea.maximumdirect.net/eric/scriptorium/internal/defaults" "gitea.maximumdirect.net/eric/scriptorium/internal/domain" "gitea.maximumdirect.net/eric/scriptorium/internal/llm" "gitea.maximumdirect.net/eric/scriptorium/internal/profile" @@ -102,13 +104,16 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes if err != nil { return nil, fmt.Errorf("%w: %w", ErrProfileLoad, err) } - effectiveModel := mergeExecutionTarget(executionProfileToTarget(execProfile), req.Execution) + effectiveModel := resolveExecutionTarget(execProfile, req.Execution) if strings.TrimSpace(effectiveModel.Endpoint) == "" { return nil, fmt.Errorf("%w: execution endpoint is required", ErrInvalidRequest) } if strings.TrimSpace(effectiveModel.Model) == "" { return nil, fmt.Errorf("%w: execution model is required", ErrInvalidRequest) } + if err := validateAPIKeyEnv(effectiveModel.APIKeyEnv); err != nil { + return nil, fmt.Errorf("%w: %w", ErrInvalidRequest, err) + } effectiveContract := resolveOutputContract(def, req.Validation) resolvedInputs := make(map[string]*domain.Artifact, len(req.Inputs)) @@ -231,11 +236,7 @@ func (r *Runner) shouldAttemptRepair(contract domain.OutputContract, validationR return contract.ValidationMode == domain.ValidationJSON || contract.ValidationMode == domain.ValidationJSONSchema } -func mergeExecutionTarget(base domain.ExecutionTarget, override *domain.ExecutionTarget) domain.ExecutionTarget { - if override == nil { - return base - } - +func mergeExecutionTarget(base domain.ExecutionTarget, override domain.ExecutionTarget) domain.ExecutionTarget { out := base if override.Endpoint != "" { out.Endpoint = override.Endpoint @@ -271,6 +272,26 @@ func mergeExecutionTarget(base domain.ExecutionTarget, override *domain.Executio return out } +func resolveExecutionTarget(profileValue *domain.ExecutionProfile, override *domain.ExecutionTarget) domain.ExecutionTarget { + out := defaults.ExecutionTargetDefault() + out = mergeExecutionTarget(out, executionProfileToTarget(profileValue)) + if override != nil { + out = mergeExecutionTarget(out, *override) + } + return out +} + +func validateAPIKeyEnv(apiKeyEnv string) error { + envName := strings.TrimSpace(apiKeyEnv) + if envName == "" { + return nil + } + if strings.TrimSpace(os.Getenv(envName)) == "" { + return fmt.Errorf("api key environment variable %q is not set", envName) + } + return nil +} + func executionProfileToTarget(p *domain.ExecutionProfile) domain.ExecutionTarget { if p == nil { return domain.ExecutionTarget{} @@ -325,16 +346,16 @@ func buildOutputArtifact(content string, format domain.OutputFormat) domain.Arti body := []byte(content) hash := sha256.Sum256(body) - contentType := "text/plain" + contentType := defaults.ContentTypeTextPlain switch format { case domain.FormatMarkdown: - contentType = "text/markdown" + contentType = defaults.ContentTypeTextMarkdown case domain.FormatJSON: - contentType = "application/json" + contentType = defaults.ContentTypeApplicationJSON } return domain.Artifact{ - Name: "output", + Name: defaults.OutputArtifactName, ContentType: contentType, Body: body, Size: int64(len(body)), diff --git a/internal/usecase/runner_test.go b/internal/usecase/runner_test.go index 6757a4b..1dc37fe 100644 --- a/internal/usecase/runner_test.go +++ b/internal/usecase/runner_test.go @@ -5,7 +5,9 @@ import ( "crypto/sha256" "encoding/hex" "errors" + "fmt" "regexp" + "strings" "testing" "gitea.maximumdirect.net/eric/scriptorium/internal/domain" @@ -20,9 +22,9 @@ type fakePromptRepo struct { } type fakeExecutionProfileRepo struct { - profile *domain.ExecutionProfile - err error - lastID string + profiles map[string]*domain.ExecutionProfile + err error + lastID string } func (f *fakeExecutionProfileRepo) GetProfile(ctx context.Context, id string) (*domain.ExecutionProfile, error) { @@ -30,10 +32,11 @@ func (f *fakeExecutionProfileRepo) GetProfile(ctx context.Context, id string) (* if f.err != nil { return nil, f.err } - if f.profile == nil { - return nil, errors.New("profile not found") + if p, ok := f.profiles[id]; ok { + cp := *p + return &cp, nil } - return f.profile, nil + return nil, errors.New("profile not found") } func (f *fakePromptRepo) GetPromptDefinition(ctx context.Context, id string, version string) (*domain.PromptDefinition, error) { @@ -103,10 +106,12 @@ type fakeRepairer struct { responses []*domain.GenerateResponse err error calls int + reqs []RepairRequest } func (f *fakeRepairer) Repair(ctx context.Context, req RepairRequest) (*domain.GenerateResponse, error) { f.calls++ + f.reqs = append(f.reqs, req) if f.err != nil { return nil, f.err } @@ -121,7 +126,8 @@ func (f *fakeRepairer) Repair(ctx context.Context, req RepairRequest) (*domain.G } func TestRunnerRunSuccessful(t *testing.T) { - repo := &fakePromptRepo{def: promptDef(domain.FormatMarkdown, domain.ValidationBasic, 0)} + promptRepo := &fakePromptRepo{def: promptDef(domain.FormatMarkdown, domain.ValidationBasic, 0)} + execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}} reader := &fakeArtifactReader{artifactsByURI: map[string]*domain.Artifact{ "a://t": {Body: []byte("transcript"), Hash: hashString("transcript")}, "a://g": {Body: []byte("glossary"), Hash: hashString("glossary")}, @@ -129,7 +135,7 @@ func TestRunnerRunSuccessful(t *testing.T) { renderer := &fakeRenderer{rendered: &domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "system", Content: "sys"}, {Role: "user", Content: "usr"}}}} llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "# recap", Usage: domain.TokenUsage{TotalTokens: 7}}} - runner := NewRunner(repo, &fakeExecutionProfileRepo{profile: defaultExecutionProfile()}, reader, renderer, llmClient, nil) + runner := NewRunner(promptRepo, execRepo, reader, renderer, llmClient, nil) res, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", PromptVersion: "1", @@ -169,36 +175,222 @@ func TestRunnerRunSuccessful(t *testing.T) { } } +func TestRunnerRunExplicitProfileIDIsUsed(t *testing.T) { + promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)} + promptRepo.def.DefaultProfile = "default-prof" + execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ + "explicit-prof": {ID: "explicit-prof", Endpoint: "http://explicit/v1", Model: "explicit"}, + "default-prof": {ID: "default-prof", Endpoint: "http://default/v1", Model: "default"}, + }} + + runner := newMinimalRunner(promptRepo, execRepo) + res, err := runner.Run(context.Background(), domain.RunRequest{ + PromptID: "p", + ProfileID: "explicit-prof", + Inputs: singleInputRef(), + }) + if err != nil { + t.Fatalf("expected no error, got %v", err) + } + if execRepo.lastID != "explicit-prof" { + t.Fatalf("expected explicit profile lookup, got %q", execRepo.lastID) + } + if res.SelectedProfileID != "explicit-prof" { + t.Fatalf("expected selected profile explicit-prof, got %q", res.SelectedProfileID) + } +} + +func TestRunnerRunPromptDefaultProfileIsUsedWhenNoExplicitProfileID(t *testing.T) { + promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)} + promptRepo.def.DefaultProfile = "from-prompt" + execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ + "from-prompt": {ID: "from-prompt", Endpoint: "http://llm/v1", Model: "m"}, + }} + + runner := newMinimalRunner(promptRepo, execRepo) + res, err := runner.Run(context.Background(), domain.RunRequest{ + PromptID: "p", + Inputs: singleInputRef(), + }) + if err != nil { + t.Fatalf("expected no error, got %v", err) + } + if execRepo.lastID != "from-prompt" { + t.Fatalf("expected prompt default profile lookup, got %q", execRepo.lastID) + } + if res.SelectedProfileID != "from-prompt" { + t.Fatalf("expected selected profile from-prompt, got %q", res.SelectedProfileID) + } +} + +func TestRunnerRunMissingExplicitProfileAndMissingDefaultProfileFails(t *testing.T) { + repo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)} + repo.def.DefaultProfile = "" + runner := newMinimalRunner(repo, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}}) + _, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p", Inputs: singleInputRef()}) + if !errors.Is(err, ErrInvalidRequest) { + t.Fatalf("expected ErrInvalidRequest, got %v", err) + } +} + +func TestRunnerRunInvalidDefaultProfileFails(t *testing.T) { + repo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)} + repo.def.DefaultProfile = "does-not-exist" + runner := newMinimalRunner(repo, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{}}) + _, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p", Inputs: singleInputRef()}) + if !errors.Is(err, ErrProfileLoad) { + t.Fatalf("expected ErrProfileLoad, got %v", err) + } +} + +func TestRunnerRunExecutionProfileLoadFailure(t *testing.T) { + repo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)} + runner := newMinimalRunner(repo, &fakeExecutionProfileRepo{err: errors.New("load failed")}) + _, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()}) + if !errors.Is(err, ErrProfileLoad) { + t.Fatalf("expected profile load failure, got %v", err) + } +} + +func TestRunnerRunExplicitRuntimeOverrideBeatsSelectedProfileValue(t *testing.T) { + promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)} + execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ + "exec": { + ID: "exec", + Endpoint: "http://profile/v1", + Model: "profile-model", + Temperature: 0.2, + MaxTokens: 500, + TopP: 0.9, + TimeoutSeconds: 120, + }, + }} + llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}} + runner := NewRunner(promptRepo, execRepo, defaultArtifactReader(), defaultRenderer(), llmClient, nil) + + res, err := runner.Run(context.Background(), domain.RunRequest{ + PromptID: "p", + ProfileID: "exec", + Inputs: singleInputRef(), + Execution: &domain.ExecutionTarget{ + Endpoint: "http://override/v1", + Model: "override-model", + Temperature: 0.7, + TimeoutSeconds: 30, + }, + }) + if err != nil { + t.Fatalf("expected no error, got %v", err) + } + if res.Endpoint != "http://override/v1" || res.ModelName != "override-model" { + t.Fatalf("expected endpoint/model override to win, got endpoint=%q model=%q", res.Endpoint, res.ModelName) + } + if res.EffectiveModelParams.Temperature != 0.7 || res.EffectiveModelParams.TimeoutSeconds != 30 { + t.Fatalf("expected numeric override to win, got %+v", res.EffectiveModelParams) + } + if res.EffectiveModelParams.TopP != 0.9 { + t.Fatalf("expected non-overridden profile top_p to remain, got %v", res.EffectiveModelParams.TopP) + } +} + +func TestRunnerRunSelectedProfileBeatsBuiltInDefault(t *testing.T) { + promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)} + execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ + "exec": { + ID: "exec", + Endpoint: "http://profile/v1", + Model: "profile-model", + TopP: 0.8, + TimeoutSeconds: 90, + }, + }} + llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}} + runner := NewRunner(promptRepo, execRepo, defaultArtifactReader(), defaultRenderer(), llmClient, nil) + + res, err := runner.Run(context.Background(), domain.RunRequest{ + PromptID: "p", + ProfileID: "exec", + Inputs: singleInputRef(), + }) + if err != nil { + t.Fatalf("expected no error, got %v", err) + } + if res.EffectiveModelParams.TopP != 0.8 { + t.Fatalf("expected profile top_p to beat default, got %v", res.EffectiveModelParams.TopP) + } + if res.EffectiveModelParams.TimeoutSeconds != 90 { + t.Fatalf("expected profile timeout to beat default, got %d", res.EffectiveModelParams.TimeoutSeconds) + } +} + +func TestRunnerRunAPIKeyEnvResolvesFromEnvironment(t *testing.T) { + t.Setenv("SCRIPTORIUM_TEST_API_KEY", "secret") + promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)} + execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ + "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: "SCRIPTORIUM_TEST_API_KEY"}, + }} + runner := NewRunner(promptRepo, execRepo, defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil) + + res, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()}) + if err != nil { + t.Fatalf("expected no error, got %v", err) + } + if res.EffectiveModelParams.APIKeyEnv != "SCRIPTORIUM_TEST_API_KEY" { + t.Fatalf("expected api_key_env name in effective params, got %q", res.EffectiveModelParams.APIKeyEnv) + } +} + +func TestRunnerRunAPIKeyEnvMissingEnvironmentValueFailsClearly(t *testing.T) { + promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)} + execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ + "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: "SCRIPTORIUM_MISSING_KEY"}, + }} + runner := NewRunner(promptRepo, execRepo, defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil) + + _, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()}) + if !errors.Is(err, ErrInvalidRequest) { + t.Fatalf("expected ErrInvalidRequest, got %v", err) + } + if !strings.Contains(err.Error(), "SCRIPTORIUM_MISSING_KEY") { + t.Fatalf("expected missing env name in error, got %v", err) + } +} + +func TestRunnerRunAPIKeyValueNotPresentInMetadata(t *testing.T) { + const envName = "SCRIPTORIUM_TEST_API_KEY" + const secret = "top-secret-value" + t.Setenv(envName, secret) + promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)} + execRepo := &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ + "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", APIKeyEnv: envName}, + }} + runner := NewRunner(promptRepo, execRepo, defaultArtifactReader(), defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil) + + res, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: singleInputRef()}) + if err != nil { + t.Fatalf("expected no error, got %v", err) + } + if res.EffectiveModelParams.APIKeyEnv != envName { + t.Fatalf("expected api key env name, got %q", res.EffectiveModelParams.APIKeyEnv) + } + metadataDump := fmt.Sprintf("%+v|%s|%s|%s|%s", res.EffectiveModelParams, res.Endpoint, res.ModelName, res.PromptHash, res.RenderedPromptHash) + if strings.Contains(metadataDump, secret) { + t.Fatalf("unexpected api key value in metadata dump: %s", metadataDump) + } +} + func TestRunnerRunPromptLoadFailure(t *testing.T) { - runner := NewRunner(&fakePromptRepo{err: errors.New("boom")}, &fakeExecutionProfileRepo{profile: defaultExecutionProfile()}, &fakeArtifactReader{}, &fakeRenderer{}, &fakeLLM{}, nil) + runner := NewRunner(&fakePromptRepo{err: errors.New("boom")}, &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}}, defaultArtifactReader(), defaultRenderer(), &fakeLLM{}, nil) _, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p"}) if !errors.Is(err, ErrProfileLoad) { t.Fatalf("expected ErrProfileLoad, got %v", err) } } -func TestRunnerRunMissingProfileSelection(t *testing.T) { - repo := &fakePromptRepo{def: &domain.PromptDefinition{ID: "p", Version: "1", Inputs: []domain.PromptInput{{Name: "transcript", Required: true}}, Templates: []domain.PromptMessageTemplate{{Role: "user", Content: "x"}}, OutputFormat: domain.FormatText, Validation: domain.OutputContract{ValidationMode: domain.ValidationNone}}} - runner := NewRunner(repo, &fakeExecutionProfileRepo{profile: defaultExecutionProfile()}, &fakeArtifactReader{}, &fakeRenderer{}, &fakeLLM{}, nil) - _, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p"}) - if !errors.Is(err, ErrInvalidRequest) { - t.Fatalf("expected invalid request, got %v", err) - } -} - -func TestRunnerRunExecutionProfileLoadFailure(t *testing.T) { - repo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)} - runner := NewRunner(repo, &fakeExecutionProfileRepo{err: errors.New("load failed")}, &fakeArtifactReader{}, &fakeRenderer{}, &fakeLLM{}, nil) - _, err := runner.Run(context.Background(), domain.RunRequest{PromptID: "p", ProfileID: "exec", Inputs: map[string]domain.ArtifactRef{"transcript": {Type: domain.ArtifactRefInline, Body: "x"}}}) - if !errors.Is(err, ErrProfileLoad) { - t.Fatalf("expected profile load failure, got %v", err) - } -} - func TestRunnerRunArtifactLoadFailure(t *testing.T) { runner := NewRunner( &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)}, - &fakeExecutionProfileRepo{profile: defaultExecutionProfile()}, + &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}}, &fakeArtifactReader{errByURI: map[string]error{"a://bad": errors.New("read failed")}}, &fakeRenderer{rendered: &domain.RenderedPrompt{}}, &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, @@ -218,8 +410,8 @@ func TestRunnerRunArtifactLoadFailure(t *testing.T) { func TestRunnerRunPromptRenderFailure(t *testing.T) { runner := NewRunner( &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)}, - &fakeExecutionProfileRepo{profile: defaultExecutionProfile()}, - &fakeArtifactReader{artifactsByURI: map[string]*domain.Artifact{"a://ok": {Body: []byte("x"), Hash: hashString("x")}}}, + &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}}, + defaultArtifactReader(), &fakeRenderer{err: errors.New("render failed")}, &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, nil, @@ -227,7 +419,7 @@ func TestRunnerRunPromptRenderFailure(t *testing.T) { _, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", ProfileID: "exec", - Inputs: map[string]domain.ArtifactRef{"transcript": {Type: domain.ArtifactRefFile, URI: "a://ok"}}, + Inputs: singleInputRef(), }) if !errors.Is(err, ErrPromptRender) { t.Fatalf("expected ErrPromptRender, got %v", err) @@ -237,36 +429,36 @@ func TestRunnerRunPromptRenderFailure(t *testing.T) { func TestRunnerRunLLMFailure(t *testing.T) { runner := NewRunner( &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)}, - &fakeExecutionProfileRepo{profile: defaultExecutionProfile()}, - &fakeArtifactReader{artifactsByURI: map[string]*domain.Artifact{"a://ok": {Body: []byte("x"), Hash: hashString("x")}}}, - &fakeRenderer{rendered: &domain.RenderedPrompt{}}, + &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}}, + defaultArtifactReader(), + defaultRenderer(), &fakeLLM{err: errors.New("llm failed")}, nil, ) _, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", ProfileID: "exec", - Inputs: map[string]domain.ArtifactRef{"transcript": {Type: domain.ArtifactRefFile, URI: "a://ok"}}, + Inputs: singleInputRef(), }) if !errors.Is(err, ErrLLMGenerate) { t.Fatalf("expected ErrLLMGenerate, got %v", err) } } -func TestRunnerRunValidationFailurePreservesRawOutput(t *testing.T) { +func TestRunnerRunValidationStillWorks(t *testing.T) { validator := &fakeValidator{result: domain.ValidationResult{Status: domain.ValidationFailed, Mode: domain.ValidationBasic, Errors: []string{"bad"}, IsValid: false}} runner := NewRunner( &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationBasic, 0)}, - &fakeExecutionProfileRepo{profile: defaultExecutionProfile()}, - &fakeArtifactReader{artifactsByURI: map[string]*domain.Artifact{"a://ok": {Body: []byte("x"), Hash: hashString("x")}}}, - &fakeRenderer{rendered: &domain.RenderedPrompt{}}, + &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}}, + defaultArtifactReader(), + defaultRenderer(), &fakeLLM{resp: &domain.GenerateResponse{Content: "raw output"}}, validator, ) res, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", ProfileID: "exec", - Inputs: map[string]domain.ArtifactRef{"transcript": {Type: domain.ArtifactRefFile, URI: "a://ok"}}, + Inputs: singleInputRef(), }) if err != nil { t.Fatalf("expected no error, got %v", err) @@ -276,21 +468,26 @@ func TestRunnerRunValidationFailurePreservesRawOutput(t *testing.T) { } } -func TestRunnerRunRepairBounded(t *testing.T) { +func TestRunnerRunStructuredRepairRemainsBoundedAndUsesEffectiveModelSettings(t *testing.T) { repairer := &fakeRepairer{responses: []*domain.GenerateResponse{{Content: `{"broken":`}, {Content: `{"still":`}}} + llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: `{"initial":`}} + runner := NewRunnerWithRepairer( &fakePromptRepo{def: promptDef(domain.FormatJSON, domain.ValidationJSON, 1)}, - &fakeExecutionProfileRepo{profile: defaultExecutionProfile()}, - &fakeArtifactReader{artifactsByURI: map[string]*domain.Artifact{"a://ok": {Body: []byte("x"), Hash: hashString("x")}}}, - &fakeRenderer{rendered: &domain.RenderedPrompt{}}, - &fakeLLM{resp: &domain.GenerateResponse{Content: `{"initial":`}}, + &fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{ + "exec": {ID: "exec", Endpoint: "http://profile/v1", Model: "profile-model", TimeoutSeconds: 55}, + }}, + defaultArtifactReader(), + defaultRenderer(), + llmClient, validate.NewStandardValidator("."), repairer, ) res, err := runner.Run(context.Background(), domain.RunRequest{ PromptID: "p", ProfileID: "exec", - Inputs: map[string]domain.ArtifactRef{"transcript": {Type: domain.ArtifactRefFile, URI: "a://ok"}}, + Inputs: singleInputRef(), + Execution: &domain.ExecutionTarget{Endpoint: "http://override/v1", Model: "override-model", TimeoutSeconds: 22}, }) if err != nil { t.Fatalf("expected no error, got %v", err) @@ -298,6 +495,15 @@ func TestRunnerRunRepairBounded(t *testing.T) { if repairer.calls != 1 || res.Validation.RepairAttempts != 1 { t.Fatalf("expected one bounded repair, calls=%d attempts=%d", repairer.calls, res.Validation.RepairAttempts) } + if len(repairer.reqs) != 1 { + t.Fatalf("expected one repair request, got %d", len(repairer.reqs)) + } + if repairer.reqs[0].Target.Endpoint != "http://override/v1" || repairer.reqs[0].Target.Model != "override-model" { + t.Fatalf("expected repair to use effective target, got %+v", repairer.reqs[0].Target) + } + if repairer.reqs[0].Target.TimeoutSeconds != 22 { + t.Fatalf("expected repair to use effective timeout, got %d", repairer.reqs[0].Target.TimeoutSeconds) + } } func promptDef(format domain.OutputFormat, mode domain.ValidationMode, attempts int) *domain.PromptDefinition { @@ -328,3 +534,28 @@ func defaultExecutionProfile() *domain.ExecutionProfile { Model: "model-from-profile", } } + +func defaultArtifactReader() *fakeArtifactReader { + return &fakeArtifactReader{artifactsByURI: map[string]*domain.Artifact{ + "a://ok": {Body: []byte("x"), Hash: hashString("x")}, + }} +} + +func defaultRenderer() *fakeRenderer { + return &fakeRenderer{rendered: &domain.RenderedPrompt{Messages: []domain.RenderedMessage{{Role: "user", Content: "hello"}}}} +} + +func singleInputRef() map[string]domain.ArtifactRef { + return map[string]domain.ArtifactRef{"transcript": {Type: domain.ArtifactRefFile, URI: "a://ok"}} +} + +func newMinimalRunner(promptRepo *fakePromptRepo, execRepo *fakeExecutionProfileRepo) *Runner { + return NewRunner( + promptRepo, + execRepo, + defaultArtifactReader(), + defaultRenderer(), + &fakeLLM{resp: &domain.GenerateResponse{Content: "ok"}}, + nil, + ) +}