Preserve comparison failures during cancellation

This commit is contained in:
2026-08-13 03:26:11 +00:00
parent 79cba800ee
commit 5e492cf1fb
10 changed files with 167 additions and 15 deletions

View File

@@ -24,6 +24,7 @@ type Config struct {
}
// Adapter owns one Promptkit engine and its opaque prepared execution handles.
// It supports concurrent Execute calls on the shared executor.
type Adapter struct {
engine *promptkit.Engine
}

View File

@@ -23,6 +23,7 @@ type fakeClient struct {
calls int
requests []promptkit.GenerateRequest
block bool
started chan struct{}
}
type recordingReader struct {
@@ -45,9 +46,13 @@ func (client *fakeClient) Generate(ctx context.Context, request promptkit.Genera
client.calls++
client.requests = append(client.requests, request)
block := client.block
started := client.started
response := client.response
err := client.err
client.mu.Unlock()
if started != nil {
started <- struct{}{}
}
if block {
<-ctx.Done()
return nil, ctx.Err()
@@ -55,6 +60,36 @@ func (client *fakeClient) Generate(ctx context.Context, request promptkit.Genera
return response, err
}
func TestExecuteSupportsConcurrentCalls(t *testing.T) {
client := &fakeClient{response: validResponse(), block: true, started: make(chan struct{}, 2)}
adapter := newTestAdapter(t, client)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
executionErrors := make(chan error, 2)
for range 2 {
go func() {
_, err := adapter.Execute(ctx, testExecuteRequest(), nil)
executionErrors <- err
}()
}
for range 2 {
select {
case <-client.started:
case <-time.After(5 * time.Second):
t.Fatal("timed out waiting for concurrent Promptkit calls")
}
}
cancel()
for range 2 {
if err := <-executionErrors; promptexec.CategoryOf(err) != promptexec.Canceled {
t.Fatalf("Execute() error/category = %v/%q", err, promptexec.CategoryOf(err))
}
}
if client.callCount() != 2 {
t.Fatalf("provider calls = %d, want 2", client.callCount())
}
}
func (client *fakeClient) callCount() int {
client.mu.Lock()
defer client.mu.Unlock()

View File

@@ -35,11 +35,21 @@ type comparisonProfileOutcome struct {
Markdown []byte
LLMDebugPath string
Error *comparison.SafeError
canceled bool
}
type comparisonProfileExecutionState uint8
const (
comparisonProfilePending comparisonProfileExecutionState = iota
comparisonProfileRunning
comparisonProfileComplete
)
func executeComparisonProfiles(ctx context.Context, req comparisonExecutionRequest) comparisonExecutionResult {
profiles := req.Inspection.Profiles
result := comparisonExecutionResult{Outcomes: make([]comparisonProfileOutcome, len(profiles))}
states := make([]comparisonProfileExecutionState, len(profiles))
for index, profile := range profiles {
result.Outcomes[index] = comparisonProfileOutcome{
Position: index + 1,
@@ -54,21 +64,22 @@ func executeComparisonProfiles(ctx context.Context, req comparisonExecutionReque
for index, profile := range profiles {
if err := ctx.Err(); err != nil {
result.Canceled = true
markUnstartedComparisonOutcomes(result.Outcomes[index:], err)
break
}
index, profile := index, profile
states[index] = comparisonProfileRunning
waitGroup.Add(1)
go func() {
defer waitGroup.Done()
result.Outcomes[index] = executeComparisonProfile(ctx, req, index, profile)
states[index] = comparisonProfileComplete
}()
}
waitGroup.Wait()
if err := ctx.Err(); err != nil {
result.Canceled = true
for index := range result.Outcomes {
if result.Outcomes[index].Status != comparison.StatusSucceeded {
if states[index] != comparisonProfileComplete || result.Outcomes[index].canceled {
markCanceledComparisonOutcome(&result.Outcomes[index], err)
}
}
@@ -99,6 +110,7 @@ func executeComparisonProfile(ctx context.Context, req comparisonExecutionReques
outcome.ValidationStatus = execution.ValidationStatus
outcome.LLMDebugPath = execution.LLMDebugPath
if err != nil {
outcome.canceled = cancellationError(err)
safe := comparisonSafeExecutionError(err)
outcome.Error = &safe
return outcome
@@ -119,12 +131,6 @@ func comparisonDebugRunID(comparisonID string, position, profileCount int, profi
return fmt.Sprintf("%s_%0*d-%s", comparisonID, comparison.OrdinalWidth(profileCount), position, comparison.ProfileSlug(profileID))
}
func markUnstartedComparisonOutcomes(outcomes []comparisonProfileOutcome, err error) {
for index := range outcomes {
markCanceledComparisonOutcome(&outcomes[index], err)
}
}
func markCanceledComparisonOutcome(outcome *comparisonProfileOutcome, err error) {
outcome.Status = comparison.StatusFailed
outcome.ValidationStatus = promptexec.ValidationSkipped
@@ -134,6 +140,12 @@ func markCanceledComparisonOutcome(outcome *comparisonProfileOutcome, err error)
outcome.Error = &safe
}
func cancellationError(err error) bool {
category := promptexec.CategoryOf(err)
return errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) ||
category == promptexec.Canceled || category == promptexec.DeadlineExceeded
}
func comparisonSafeExecutionError(err error) comparison.SafeError {
category := promptexec.CategoryOf(err)
if category == "" {

View File

@@ -7,6 +7,7 @@ import (
"os"
"path/filepath"
"strings"
"sync"
"testing"
"time"
@@ -307,6 +308,57 @@ func TestCompareDetailedCancellationPreservesPublishedBundle(t *testing.T) {
}
}
func TestCompareDetailedPreservesCompletedProfileFailureWhenCanceled(t *testing.T) {
bundle := generationBundle(t)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
failureStarted := make(chan struct{})
var signalFailure sync.Once
executor := &generationExecutor{
validations: map[string]promptexec.ValidationStatus{"weather-light": promptexec.ValidationFailed},
waitForCancellation: map[string]bool{"weather-deep": true},
beforeExecute: func(request promptexec.ExecuteRequest) {
if request.ProfileID == "weather-light" {
signalFailure.Do(func() { close(failureStarted) })
}
},
}
results := make(chan struct {
result *ComparisonResult
err error
}, 1)
go func() {
result, err := CompareDetailed(ctx, ComparisonRequest{
Config: comparisonConfig(), Report: ReportDaily, ProfileIDs: []string{"weather-light", "weather-deep"},
WorkingDir: t.TempDir(), Date: generationTime("2026-05-29T12:00:00-05:00"),
Clock: timeutil.FixedClock{Time: generationTime("2026-05-29T08:30:00-05:00")},
Collector: &generationCollector{bundle: &bundle}, Executor: executor,
})
results <- struct {
result *ComparisonResult
err error
}{result: result, err: err}
}()
select {
case <-failureStarted:
cancel()
case <-time.After(5 * time.Second):
t.Fatal("timed out waiting for the completed profile failure")
}
completed := <-results
if !errors.Is(completed.err, context.Canceled) || completed.result == nil || completed.result.ManifestPath != "" || completed.result.DataPackagePath != "" || completed.result.Succeeded != 0 || completed.result.Failed != 2 {
t.Fatalf("CompareDetailed() result/error = %#v/%v", completed.result, completed.err)
}
failed, canceled := completed.result.Results[0], completed.result.Results[1]
if failed.Error == nil || failed.Error.Category != string(promptexec.ValidationRejected) || failed.ValidationStatus != promptexec.ValidationFailed || failed.ReportPath != "" {
t.Fatalf("completed failure = %#v", failed)
}
if canceled.Error == nil || canceled.Error.Category != string(promptexec.Canceled) || canceled.ValidationStatus != promptexec.ValidationSkipped || canceled.ReportPath != "" {
t.Fatalf("canceled profile = %#v", canceled)
}
}
func TestCompareDetailedLeavesExistingBundleWhenPublicationPreflightChanges(t *testing.T) {
workingDir := t.TempDir()
target := filepath.Join(workingDir, "comparison-output")

View File

@@ -66,7 +66,9 @@ type generationExecutor struct {
beforeExecute func(promptexec.ExecuteRequest)
cancelBeforeReturn context.CancelFunc
validation promptexec.ValidationStatus
validations map[string]promptexec.ValidationStatus
rawOutput []byte
waitForCancellation map[string]bool
failedPrompt string
skipPreparation bool
preparationCalls int
@@ -95,7 +97,7 @@ func (e *generationExecutor) InspectProfile(_ context.Context, id string) (promp
}
return promptexec.ProfileInspection{ProfileID: id, BackendID: "fixture", ModelName: "fixture-model"}, nil
}
func (e *generationExecutor) Execute(_ context.Context, req promptexec.ExecuteRequest, callback promptexec.PreparationCallback) (*promptexec.Execution, error) {
func (e *generationExecutor) Execute(ctx context.Context, req promptexec.ExecuteRequest, callback promptexec.PreparationCallback) (*promptexec.Execution, error) {
stamp := time.Date(2026, 5, 29, 15, 0, 0, 0, time.UTC)
generationExecutorMu.Lock()
skipPreparation := e.skipPreparation
@@ -124,7 +126,11 @@ func (e *generationExecutor) Execute(_ context.Context, req promptexec.ExecuteRe
profileErr := e.executeErrors[req.ProfileID]
executeErr := e.executeErr
status := e.validation
if profileStatus, ok := e.validations[req.ProfileID]; ok {
status = profileStatus
}
rawOutput := append([]byte(nil), e.rawOutput...)
waitForCancellation := e.waitForCancellation[req.ProfileID]
failedPrompt := e.failedPrompt
cancelBeforeReturn := e.cancelBeforeReturn
complete := e.complete
@@ -132,6 +138,10 @@ func (e *generationExecutor) Execute(_ context.Context, req promptexec.ExecuteRe
if beforeExecute != nil {
beforeExecute(req)
}
if waitForCancellation {
<-ctx.Done()
return nil, ctx.Err()
}
if profileErr != nil {
return nil, profileErr
}

View File

@@ -215,6 +215,40 @@ func TestCompareCommandWritesStructuredPartialFailure(t *testing.T) {
}
}
func TestCompareCommandPreservesMixedCancellationOutcomes(t *testing.T) {
workingDir := t.TempDir()
configPath := comparisonConfigPath(t, "weather_api:\n base_url: https://weather.api.example.com/\n")
profileFailure := comparison.NewSafeError("validation_rejected", "validate prompt execution failed")
canceledProfile := comparison.NewSafeError("canceled", "profile execution canceled")
result := comparisonResult("/reports/comparison-daily", []app.ComparisonProfileResult{
{Position: 1, ProfileID: "weather-light", ModelName: "light", Status: comparison.StatusFailed, ValidationStatus: promptexec.ValidationFailed, Error: &profileFailure},
{Position: 2, ProfileID: "weather-deep", ModelName: "deep", Status: comparison.StatusFailed, ValidationStatus: promptexec.ValidationSkipped, Error: &canceledProfile},
})
runner := comparisonRunner(t, workingDir)
runner.compareDetailed = func(context.Context, app.ComparisonRequest) (*app.ComparisonResult, error) {
return result, context.Canceled
}
var stdout, stderr bytes.Buffer
err := runner.Run(context.Background(), []string{"compare", "daily", "--date", "2026-05-29", "--profile", "weather-light", "--profile", "weather-deep", "--config", configPath}, &stdout, &stderr)
if !errors.Is(err, context.Canceled) || stderr.Len() != 0 {
t.Fatalf("Run() error/stderr = %v/%q", err, stderr.String())
}
var summary comparisonSummary
if err := json.Unmarshal(stdout.Bytes(), &summary); err != nil {
t.Fatalf("decode summary: %v\n%s", err, stdout.String())
}
if summary.Status != summaryStatusFailed || summary.Error == nil || summary.Error.Category != "canceled" || len(summary.Results) != 2 {
t.Fatalf("summary = %#v", summary)
}
if first := summary.Results[0]; first.Error == nil || first.Error.Category != "validation_rejected" || first.ValidationStatus != string(promptexec.ValidationFailed) {
t.Fatalf("completed profile summary = %#v", first)
}
if second := summary.Results[1]; second.Error == nil || second.Error.Category != "canceled" || second.ValidationStatus != string(promptexec.ValidationSkipped) {
t.Fatalf("canceled profile summary = %#v", second)
}
}
func TestCompareCommandReportsCommittedBundleWhenCleanupFails(t *testing.T) {
workingDir := t.TempDir()
configPath := comparisonConfigPath(t, "weather_api:\n base_url: https://weather.api.example.com/\n")

View File

@@ -20,6 +20,8 @@ const (
// error, Execute must not call the provider. Completed validation rejection is
// returned as an Execution with a failed Validation status; operational failures
// return no Execution. Sensitive debug values are populated only when requested.
// Comparison may invoke Execute concurrently on one shared Executor, so every
// implementation must support concurrent calls.
type Executor interface {
InspectPrompt(context.Context, string, string) (PromptInspection, error)
InspectProfile(context.Context, string) (ProfileInspection, error)