From 033bc93d3cca49349bc2a44657fda8bc24f07131 Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Tue, 28 Jul 2026 00:45:59 +0000 Subject: [PATCH] Use the public engine for CLI run and render --- internal/adapter/cli/run.go | 48 +++++++++++++++---------- internal/adapter/cli/run_test.go | 52 +++++++++++++++++++++++----- internal/format/prepared_run.go | 16 ++++----- internal/format/prepared_run_test.go | 36 +++++++++---------- 4 files changed, 97 insertions(+), 55 deletions(-) diff --git a/internal/adapter/cli/run.go b/internal/adapter/cli/run.go index 060b55a..37c9e12 100644 --- a/internal/adapter/cli/run.go +++ b/internal/adapter/cli/run.go @@ -12,11 +12,11 @@ import ( "strings" "time" + "gitea.maximumdirect.net/eric/scriptorium" httpadapter "gitea.maximumdirect.net/eric/scriptorium/internal/adapter/http" artifactadapter "gitea.maximumdirect.net/eric/scriptorium/internal/artifact" appconfig "gitea.maximumdirect.net/eric/scriptorium/internal/config" "gitea.maximumdirect.net/eric/scriptorium/internal/defaults" - "gitea.maximumdirect.net/eric/scriptorium/internal/domain" renderformat "gitea.maximumdirect.net/eric/scriptorium/internal/format" "gitea.maximumdirect.net/eric/scriptorium/internal/llm" "gitea.maximumdirect.net/eric/scriptorium/internal/profile/builtin" @@ -140,15 +140,13 @@ func runCommand(args []string, stdout, stderr io.Writer) int { return ExitRuntimeError } - llmClient, err := newOpenAIClient() + engine, err := newEngine(cfg) if err != nil { - fmt.Fprintf(stderr, "llm client error: %v\n", err) + fmt.Fprintf(stderr, "engine error: %v\n", err) return ExitRuntimeError } - runner := newRunner(cfg.promptDir, cfg.profileDir, cfg.schemaDir, llmClient) - - res, runErr := runner.Run(context.Background(), req) + res, runErr := engine.Run(context.Background(), req) if runErr != nil { fmt.Fprintf(stderr, "run error: %v\n", runErr) return ExitRuntimeError @@ -176,9 +174,13 @@ func renderCommand(args []string, stdout, stderr io.Writer) int { return ExitRuntimeError } - runner := newRunner(cfg.promptDir, cfg.profileDir, cfg.schemaDir, nil) + engine, err := newEngine(&cfg.runConfig) + if err != nil { + fmt.Fprintf(stderr, "engine error: %v\n", err) + return ExitRuntimeError + } - prepared, prepErr := runner.Prepare(context.Background(), req) + prepared, prepErr := engine.Prepare(context.Background(), req) if prepErr != nil { fmt.Fprintf(stderr, "render error: %v\n", prepErr) return ExitRuntimeError @@ -565,28 +567,36 @@ func newOpenAIClient() (*llm.OpenAICompatibleClient, error) { }) } -func buildRunRequestFromConfig(cfg *runConfig) (domain.RunRequest, error) { +func newEngine(cfg *runConfig, options ...scriptorium.Option) (*scriptorium.Engine, error) { + return scriptorium.NewEngine(scriptorium.Config{ + PromptDir: cfg.promptDir, + ProfileDir: cfg.profileDir, + SchemaDir: cfg.schemaDir, + }, options...) +} + +func buildRunRequestFromConfig(cfg *runConfig) (scriptorium.RunRequest, error) { inputMappings, err := parseMappings(cfg.inputRaw, false) if err != nil { - return domain.RunRequest{}, fmt.Errorf("input parse error: %w", err) + return scriptorium.RunRequest{}, fmt.Errorf("input parse error: %w", err) } varMappings := map[string]string{} if len(cfg.varRaw) > 0 { varMappings, err = parseMappings(cfg.varRaw, false) if err != nil { - return domain.RunRequest{}, fmt.Errorf("var parse error: %w", err) + return scriptorium.RunRequest{}, fmt.Errorf("var parse error: %w", err) } } - inputs := make(map[string]domain.ArtifactRef, len(inputMappings)) + inputs := make(map[string]scriptorium.ArtifactRef, len(inputMappings)) for name, path := range inputMappings { - inputs[name] = domain.ArtifactRef{Type: domain.ArtifactRefFile, URI: path} + inputs[name] = scriptorium.File(path) } - var modelOverride *domain.ExecutionTargetOverride + var modelOverride *scriptorium.ExecutionTargetOverride if cfg.llmBaseURLSet || cfg.modelSet || cfg.temperatureSet || cfg.maxTokensSet || cfg.topPSet || cfg.apiKeyEnvSet || cfg.timeoutSet { - modelOverride = &domain.ExecutionTargetOverride{ + modelOverride = &scriptorium.ExecutionTargetOverride{ Endpoint: cfg.llmBaseURL, Model: cfg.model, APIKeyEnv: cfg.apiKeyEnv, @@ -606,7 +616,7 @@ func buildRunRequestFromConfig(cfg *runConfig) (domain.RunRequest, error) { } } - return domain.RunRequest{ + return scriptorium.RunRequest{ PromptID: cfg.promptID, ProfileID: cfg.profileID, Inputs: inputs, @@ -670,17 +680,17 @@ func writeOutput(stdout io.Writer, outputPath string, body []byte) error { return os.WriteFile(outputPath, body, 0644) } -func determineExitCode(runErr error, result *domain.RunResult) int { +func determineExitCode(runErr error, result *scriptorium.RunResult) int { if runErr != nil { return ExitRuntimeError } - if result != nil && result.Validation.Status == domain.ValidationFailed { + if result != nil && result.Validation.Status == scriptorium.ValidationFailed { return ExitValidationFailed } return ExitOK } -func printSummary(stderr io.Writer, res *domain.RunResult) { +func printSummary(stderr io.Writer, res *scriptorium.RunResult) { if res == nil { return } diff --git a/internal/adapter/cli/run_test.go b/internal/adapter/cli/run_test.go index f2400e0..a8f3430 100644 --- a/internal/adapter/cli/run_test.go +++ b/internal/adapter/cli/run_test.go @@ -17,9 +17,9 @@ import ( "testing" "time" + "gitea.maximumdirect.net/eric/scriptorium" appconfig "gitea.maximumdirect.net/eric/scriptorium/internal/config" "gitea.maximumdirect.net/eric/scriptorium/internal/defaults" - "gitea.maximumdirect.net/eric/scriptorium/internal/domain" renderformat "gitea.maximumdirect.net/eric/scriptorium/internal/format" ) @@ -655,6 +655,40 @@ func TestRunAndRenderBuildEquivalentRuntimeOverrideRequestsForSharedFlags(t *tes } } +func TestBuildRunRequestPreservesNumericOverridePresence(t *testing.T) { + omitted, err := buildRunRequestFromConfig(&runConfig{ + promptID: "prompt-1", + inputRaw: []string{"transcript=./transcript.md"}, + }) + if err != nil { + t.Fatalf("expected omitted override request to build, got %v", err) + } + if omitted.Execution != nil { + t.Fatalf("expected omitted numeric flags to leave execution override nil, got %#v", omitted.Execution) + } + + explicitZeros, err := buildRunRequestFromConfig(&runConfig{ + promptID: "prompt-1", + inputRaw: []string{"transcript=./transcript.md"}, + temperatureSet: true, + maxTokensSet: true, + topPSet: true, + timeoutSet: true, + }) + if err != nil { + t.Fatalf("expected explicit zero override request to build, got %v", err) + } + if explicitZeros.Execution == nil { + t.Fatal("expected explicit numeric flags to create execution override") + } + if explicitZeros.Execution.Temperature == nil || explicitZeros.Execution.MaxTokens == nil || explicitZeros.Execution.TopP == nil || explicitZeros.Execution.TimeoutSeconds == nil { + t.Fatalf("expected explicit zero numeric overrides to remain non-nil, got %#v", explicitZeros.Execution) + } + if *explicitZeros.Execution.Temperature != 0 || *explicitZeros.Execution.MaxTokens != 0 || *explicitZeros.Execution.TopP != 0 || *explicitZeros.Execution.TimeoutSeconds != 0 { + t.Fatalf("expected explicit numeric overrides to retain zero values, got %#v", explicitZeros.Execution) + } +} + func TestParseRunArgsFailsClearlyWhenNoEffectivePromptDir(t *testing.T) { configPath := writeAppConfigFile(t, ` profile_dir: ./profiles @@ -731,13 +765,13 @@ func TestDetermineExitCode(t *testing.T) { if got := determineExitCode(errors.New("boom"), nil); got != ExitRuntimeError { t.Fatalf("expected runtime exit code, got %d", got) } - if got := determineExitCode(nil, &domain.RunResult{Validation: domain.ValidationResult{Status: domain.ValidationFailed}}); got != ExitValidationFailed { + if got := determineExitCode(nil, &scriptorium.RunResult{Validation: scriptorium.ValidationResult{Status: scriptorium.ValidationFailed}}); got != ExitValidationFailed { t.Fatalf("expected validation exit code, got %d", got) } - if got := determineExitCode(nil, &domain.RunResult{Validation: domain.ValidationResult{Status: domain.ValidationPassed}}); got != ExitOK { + if got := determineExitCode(nil, &scriptorium.RunResult{Validation: scriptorium.ValidationResult{Status: scriptorium.ValidationPassed}}); got != ExitOK { t.Fatalf("expected success exit code for passed validation, got %d", got) } - if got := determineExitCode(nil, &domain.RunResult{Validation: domain.ValidationResult{Status: domain.ValidationSkipped}}); got != ExitOK { + if got := determineExitCode(nil, &scriptorium.RunResult{Validation: scriptorium.ValidationResult{Status: scriptorium.ValidationSkipped}}); got != ExitOK { t.Fatalf("expected success exit code for skipped validation, got %d", got) } } @@ -1208,12 +1242,12 @@ func TestWriteOutputAndSummaryUseSeparateWriters(t *testing.T) { if err := writeOutput(&stdout, "", []byte("artifact-body")); err != nil { t.Fatalf("unexpected writeOutput error: %v", err) } - printSummary(&stderr, &domain.RunResult{ + printSummary(&stderr, &scriptorium.RunResult{ PromptID: "p", PromptVersion: "1", SelectedProfileID: "exec", ModelName: "m", - Validation: domain.ValidationResult{Status: domain.ValidationPassed, Mode: domain.ValidationBasic}, + Validation: scriptorium.ValidationResult{Status: scriptorium.ValidationPassed, Mode: scriptorium.ValidationBasic}, RenderedPromptHash: "h", InputHashes: map[string]string{"in": "x"}, }) @@ -1232,15 +1266,15 @@ func TestWriteOutputAndSummaryUseSeparateWriters(t *testing.T) { func TestPrintSummaryIncludesCacheUsageWhenPresent(t *testing.T) { var stderr bytes.Buffer - printSummary(&stderr, &domain.RunResult{ + printSummary(&stderr, &scriptorium.RunResult{ PromptID: "p", PromptVersion: "1", SelectedProfileID: "exec", ModelName: "m", - Validation: domain.ValidationResult{Status: domain.ValidationPassed, Mode: domain.ValidationBasic}, + Validation: scriptorium.ValidationResult{Status: scriptorium.ValidationPassed, Mode: scriptorium.ValidationBasic}, RenderedPromptHash: "h", InputHashes: map[string]string{"in": "x"}, - Usage: domain.TokenUsage{ + Usage: scriptorium.TokenUsage{ PromptTokens: 10, CompletionTokens: 5, TotalTokens: 15, diff --git a/internal/format/prepared_run.go b/internal/format/prepared_run.go index 904191f..fc1cb79 100644 --- a/internal/format/prepared_run.go +++ b/internal/format/prepared_run.go @@ -1,4 +1,4 @@ -// Package format formats already-prepared domain data for adapters. +// Package format formats already-prepared public data for adapters. package format import ( @@ -9,7 +9,7 @@ import ( "sort" "strings" - "gitea.maximumdirect.net/eric/scriptorium/internal/domain" + "gitea.maximumdirect.net/eric/scriptorium" ) var ErrUnknownPreparedRunFormat = errors.New("unknown prepared run format") @@ -26,7 +26,7 @@ const ( // PreparedRunFormatter serializes a prepared run without performing use case work. type PreparedRunFormatter interface { - Format(prepared *domain.PreparedRun) ([]byte, error) + Format(prepared *scriptorium.PreparedRun) ([]byte, error) } // ParsePreparedRunOutputFormat parses a format name. @@ -56,7 +56,7 @@ func FormatterForPreparedRun(outputFormat PreparedRunOutputFormat) (PreparedRunF } // FormatPreparedRun formats a prepared run using the selected format. -func FormatPreparedRun(prepared *domain.PreparedRun, outputFormat PreparedRunOutputFormat) ([]byte, error) { +func FormatPreparedRun(prepared *scriptorium.PreparedRun, outputFormat PreparedRunOutputFormat) ([]byte, error) { formatter, err := FormatterForPreparedRun(outputFormat) if err != nil { return nil, err @@ -65,7 +65,7 @@ func FormatPreparedRun(prepared *domain.PreparedRun, outputFormat PreparedRunOut } // FormatPreparedRunByName parses a format name and formats a prepared run. -func FormatPreparedRunByName(prepared *domain.PreparedRun, rawFormat string) ([]byte, error) { +func FormatPreparedRunByName(prepared *scriptorium.PreparedRun, rawFormat string) ([]byte, error) { outputFormat, err := ParsePreparedRunOutputFormat(rawFormat) if err != nil { return nil, err @@ -75,7 +75,7 @@ func FormatPreparedRunByName(prepared *domain.PreparedRun, rawFormat string) ([] type jsonPreparedRunFormatter struct{} -func (jsonPreparedRunFormatter) Format(prepared *domain.PreparedRun) ([]byte, error) { +func (jsonPreparedRunFormatter) Format(prepared *scriptorium.PreparedRun) ([]byte, error) { if prepared == nil { return nil, errors.New("prepared run is nil") } @@ -84,7 +84,7 @@ func (jsonPreparedRunFormatter) Format(prepared *domain.PreparedRun) ([]byte, er type textPreparedRunFormatter struct{} -func (textPreparedRunFormatter) Format(prepared *domain.PreparedRun) ([]byte, error) { +func (textPreparedRunFormatter) Format(prepared *scriptorium.PreparedRun) ([]byte, error) { if prepared == nil { return nil, errors.New("prepared run is nil") } @@ -146,7 +146,7 @@ func (textPreparedRunFormatter) Format(prepared *domain.PreparedRun) ([]byte, er fmt.Fprintln(&b, "messages:") roleOrder := make([]string, 0) - byRole := make(map[string][]domain.RenderedMessage) + byRole := make(map[string][]scriptorium.RenderedMessage) for _, msg := range prepared.Messages { if _, exists := byRole[msg.Role]; !exists { roleOrder = append(roleOrder, msg.Role) diff --git a/internal/format/prepared_run_test.go b/internal/format/prepared_run_test.go index 3377e07..442a5c1 100644 --- a/internal/format/prepared_run_test.go +++ b/internal/format/prepared_run_test.go @@ -6,7 +6,7 @@ import ( "strings" "testing" - "gitea.maximumdirect.net/eric/scriptorium/internal/domain" + "gitea.maximumdirect.net/eric/scriptorium" ) func TestTextFormatterIncludesPreparedRunDetails(t *testing.T) { @@ -94,9 +94,8 @@ func TestTextFormatterDoesNotIncludeResolvedAPIKeyValue(t *testing.T) { func TestTextFormatterDoesNotIncludeDirectAPIKeyValue(t *testing.T) { const directKey = "direct-format-key" + // PreparedRun intentionally has no field for direct API keys. prepared := samplePreparedRun() - prepared.EffectiveModelParams.APIKey = directKey - out, err := FormatPreparedRun(prepared, PreparedRunFormatText) if err != nil { t.Fatalf("expected no error, got %v", err) @@ -108,12 +107,12 @@ func TestTextFormatterDoesNotIncludeDirectAPIKeyValue(t *testing.T) { func TestTextFormatterIncludesMessageCacheControlBeforeContent(t *testing.T) { prepared := samplePreparedRun() - prepared.Messages = []domain.RenderedMessage{ + prepared.Messages = []scriptorium.RenderedMessage{ { Role: "system", Content: "System guidance.", - CacheControl: &domain.CacheControl{ - Type: domain.CacheControlEphemeral, + CacheControl: &scriptorium.CacheControl{ + Type: scriptorium.CacheControlEphemeral, TTL: "1h", }, }, @@ -148,12 +147,12 @@ func TestTextFormatterIncludesSessionIDWhenPresent(t *testing.T) { func TestTextFormatterOmitsEmptyCacheControlTTL(t *testing.T) { prepared := samplePreparedRun() - prepared.Messages = []domain.RenderedMessage{ + prepared.Messages = []scriptorium.RenderedMessage{ { Role: "system", Content: "System guidance.", - CacheControl: &domain.CacheControl{ - Type: domain.CacheControlEphemeral, + CacheControl: &scriptorium.CacheControl{ + Type: scriptorium.CacheControlEphemeral, }, }, } @@ -231,12 +230,12 @@ func TestJSONFormatterEmitsValidJSONAndIncludesPreparedRunFields(t *testing.T) { func TestJSONFormatterIncludesMessageCacheControlOnlyWhenPresent(t *testing.T) { prepared := samplePreparedRun() - prepared.Messages = []domain.RenderedMessage{ + prepared.Messages = []scriptorium.RenderedMessage{ { Role: "system", Content: "System guidance.", - CacheControl: &domain.CacheControl{ - Type: domain.CacheControlEphemeral, + CacheControl: &scriptorium.CacheControl{ + Type: scriptorium.CacheControlEphemeral, TTL: "1h", }, }, @@ -262,7 +261,7 @@ func TestJSONFormatterIncludesMessageCacheControlOnlyWhenPresent(t *testing.T) { if !ok { t.Fatalf("expected first message cache_control, got %#v", decoded.Messages[0]) } - if cacheControl["type"] != string(domain.CacheControlEphemeral) || cacheControl["ttl"] != "1h" { + if cacheControl["type"] != string(scriptorium.CacheControlEphemeral) || cacheControl["ttl"] != "1h" { t.Fatalf("unexpected cache_control payload: %#v", cacheControl) } if _, ok := decoded.Messages[1]["cache_control"]; ok { @@ -285,9 +284,8 @@ func TestJSONFormatterDoesNotIncludeResolvedAPIKeyValue(t *testing.T) { func TestJSONFormatterDoesNotIncludeDirectAPIKeyValue(t *testing.T) { const directKey = "direct-format-key" + // PreparedRun intentionally has no field for direct API keys. prepared := samplePreparedRun() - prepared.EffectiveModelParams.APIKey = directKey - out, err := FormatPreparedRun(prepared, PreparedRunFormatJSON) if err != nil { t.Fatalf("expected no error, got %v", err) @@ -342,13 +340,13 @@ func TestFormatPreparedRunByNameUnknownFailsClearly(t *testing.T) { } } -func samplePreparedRun() *domain.PreparedRun { - return &domain.PreparedRun{ +func samplePreparedRun() *scriptorium.PreparedRun { + return &scriptorium.PreparedRun{ PromptID: "prompt.id", PromptVersion: "v1", PromptHash: "prompt-hash", SelectedProfileID: "local-fast", - EffectiveModelParams: domain.ExecutionTarget{ + EffectiveModelParams: scriptorium.ExecutionTarget{ Endpoint: "http://llm/v1", Model: "gpt-test", Temperature: 0.4, @@ -364,7 +362,7 @@ func samplePreparedRun() *domain.PreparedRun { "glossary": "hash-glossary", }, RenderedPromptHash: "rendered-hash", - Messages: []domain.RenderedMessage{ + Messages: []scriptorium.RenderedMessage{ {Role: "system", Content: "System guidance."}, {Role: "user", Content: "Summarize the transcript.\nInclude key entities."}, {Role: "user", Content: "Second user message."},