Add direct per-run session overrides
This commit is contained in:
@@ -160,6 +160,7 @@ func (r *Runner) Run(ctx context.Context, req domain.RunRequest) (*domain.RunRes
|
||||
PromptID: prepared.PromptID,
|
||||
PromptVersion: prepared.PromptVersion,
|
||||
PromptHash: prepared.PromptHash,
|
||||
SessionID: prepared.SessionID,
|
||||
RenderedPromptHash: prepared.RenderedPromptHash,
|
||||
SelectedProfileID: prepared.SelectedProfileID,
|
||||
SelectedBackendID: prepared.SelectedBackendID,
|
||||
@@ -178,6 +179,10 @@ func (r *Runner) Prepare(ctx context.Context, req domain.RunRequest) (*domain.Pr
|
||||
if strings.TrimSpace(req.PromptID) == "" {
|
||||
return nil, fmt.Errorf("%w: prompt id is required", ErrInvalidRequest)
|
||||
}
|
||||
directSessionID, err := domain.NormalizeSessionID(req.SessionID)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: session_id: %v", ErrInvalidRequest, err)
|
||||
}
|
||||
|
||||
start := time.Now().UTC()
|
||||
|
||||
@@ -251,10 +256,19 @@ func (r *Runner) Prepare(ctx context.Context, req domain.RunRequest) (*domain.Pr
|
||||
inputHashes[name] = art.Hash
|
||||
}
|
||||
|
||||
renderedPrompt, err := r.renderer.Render(ctx, def, resolvedInputs, req.Vars)
|
||||
definitionToRender := def
|
||||
if directSessionID != "" {
|
||||
definitionCopy := *def
|
||||
definitionCopy.SessionID = ""
|
||||
definitionToRender = &definitionCopy
|
||||
}
|
||||
renderedPrompt, err := r.renderer.Render(ctx, definitionToRender, resolvedInputs, req.Vars)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %w", ErrPromptRender, err)
|
||||
}
|
||||
if directSessionID != "" {
|
||||
renderedPrompt.SessionID = directSessionID
|
||||
}
|
||||
|
||||
end := time.Now().UTC()
|
||||
return &domain.PreparedRun{
|
||||
|
||||
@@ -222,6 +222,145 @@ func TestRunnerPrepareWithExplicitProfileSelection(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestRunnerDirectSessionResolution(t *testing.T) {
|
||||
t.Run("direct value wins and changes only the rendered prompt hash", func(t *testing.T) {
|
||||
def := promptDef(domain.FormatText, domain.ValidationNone, 0)
|
||||
def.SessionID = "template-{{.template_session}}"
|
||||
promptRepo := &fakePromptRepo{def: def}
|
||||
runner := NewRunner(
|
||||
promptRepo,
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
defaultArtifactReader(),
|
||||
prompt.NewGoRenderer(),
|
||||
&fakeLLM{forbid: true},
|
||||
nil,
|
||||
)
|
||||
req := domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
Inputs: singleInputRef(),
|
||||
Vars: map[string]string{"template_session": "from-template"},
|
||||
}
|
||||
|
||||
req.SessionID = " direct-one "
|
||||
first, err := runner.Prepare(context.Background(), req)
|
||||
if err != nil {
|
||||
t.Fatalf("prepare first direct session: %v", err)
|
||||
}
|
||||
req.SessionID = "direct-two"
|
||||
second, err := runner.Prepare(context.Background(), req)
|
||||
if err != nil {
|
||||
t.Fatalf("prepare second direct session: %v", err)
|
||||
}
|
||||
|
||||
if first.SessionID != "direct-one" || second.SessionID != "direct-two" {
|
||||
t.Fatalf("direct sessions were not normalized: first=%q second=%q", first.SessionID, second.SessionID)
|
||||
}
|
||||
if first.PromptHash != second.PromptHash {
|
||||
t.Fatalf("direct session changed prompt-definition hash: first=%q second=%q", first.PromptHash, second.PromptHash)
|
||||
}
|
||||
if first.RenderedPromptHash == second.RenderedPromptHash {
|
||||
t.Fatal("changing direct session did not change rendered-prompt hash")
|
||||
}
|
||||
if def.SessionID != "template-{{.template_session}}" {
|
||||
t.Fatalf("repository-owned prompt definition was mutated: %q", def.SessionID)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("direct value bypasses failing session template without changing messages", func(t *testing.T) {
|
||||
def := promptDef(domain.FormatText, domain.ValidationNone, 0)
|
||||
def.SessionID = "{{.missing_session}}"
|
||||
def.Templates = []domain.PromptMessageTemplate{
|
||||
{Role: "user", Content: "Hello {{.name}}"},
|
||||
}
|
||||
runner := NewRunner(
|
||||
&fakePromptRepo{def: def},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
defaultArtifactReader(),
|
||||
prompt.NewGoRenderer(),
|
||||
&fakeLLM{forbid: true},
|
||||
nil,
|
||||
)
|
||||
|
||||
prepared, err := runner.Prepare(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
SessionID: "direct-session",
|
||||
Inputs: singleInputRef(),
|
||||
Vars: map[string]string{"name": "Rin"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare with direct session: %v", err)
|
||||
}
|
||||
if prepared.SessionID != "direct-session" {
|
||||
t.Fatalf("prepared session id = %q, want direct-session", prepared.SessionID)
|
||||
}
|
||||
if len(prepared.Messages) != 1 || prepared.Messages[0].Content != "Hello Rin" {
|
||||
t.Fatalf("message templates did not render normally: %+v", prepared.Messages)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("blank direct value retains prompt template behavior", func(t *testing.T) {
|
||||
def := promptDef(domain.FormatText, domain.ValidationNone, 0)
|
||||
def.SessionID = " template-{{.template_session}} "
|
||||
runner := NewRunner(
|
||||
&fakePromptRepo{def: def},
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
defaultArtifactReader(),
|
||||
prompt.NewGoRenderer(),
|
||||
&fakeLLM{forbid: true},
|
||||
nil,
|
||||
)
|
||||
|
||||
prepared, err := runner.Prepare(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
SessionID: " \t ",
|
||||
Inputs: singleInputRef(),
|
||||
Vars: map[string]string{"template_session": "rendered"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("prepare with prompt session template: %v", err)
|
||||
}
|
||||
if prepared.SessionID != "template-rendered" {
|
||||
t.Fatalf("prepared session id = %q, want template-rendered", prepared.SessionID)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("overlong direct value fails before loading or generation", func(t *testing.T) {
|
||||
promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)}
|
||||
llmClient := &fakeLLM{resp: &domain.GenerateResponse{Content: "unexpected"}}
|
||||
runner := NewRunner(
|
||||
promptRepo,
|
||||
&fakeExecutionProfileRepo{profiles: map[string]*domain.ExecutionProfile{"exec": defaultExecutionProfile()}},
|
||||
nil,
|
||||
defaultArtifactReader(),
|
||||
defaultRenderer(),
|
||||
llmClient,
|
||||
nil,
|
||||
)
|
||||
|
||||
_, err := runner.Run(context.Background(), domain.RunRequest{
|
||||
PromptID: "p",
|
||||
ProfileID: "exec",
|
||||
SessionID: strings.Repeat("界", domain.SessionIDMaxLength+1),
|
||||
Inputs: singleInputRef(),
|
||||
})
|
||||
if !errors.Is(err, ErrInvalidRequest) {
|
||||
t.Fatalf("expected ErrInvalidRequest, got %v", err)
|
||||
}
|
||||
if promptRepo.lastID != "" {
|
||||
t.Fatalf("invalid direct session loaded prompt %q", promptRepo.lastID)
|
||||
}
|
||||
if llmClient.calls != 0 {
|
||||
t.Fatalf("invalid direct session invoked generation %d times", llmClient.calls)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestRunnerPrepareUsesPromptDefaultProfileWhenNoExplicitProfileID(t *testing.T) {
|
||||
promptRepo := &fakePromptRepo{def: promptDef(domain.FormatText, domain.ValidationNone, 0)}
|
||||
promptRepo.def.DefaultProfile = "from-prompt"
|
||||
@@ -900,6 +1039,9 @@ func TestRunnerRunSuccessful(t *testing.T) {
|
||||
if llmClient.lastReq.Prompt.SessionID != "session-123" {
|
||||
t.Fatalf("expected session id to be sent to llm, got %q", llmClient.lastReq.Prompt.SessionID)
|
||||
}
|
||||
if res.SessionID != "session-123" {
|
||||
t.Fatalf("expected session id in run result, got %q", res.SessionID)
|
||||
}
|
||||
if res.Usage.TotalTokens != 7 {
|
||||
t.Fatalf("expected token usage to be retained, got %+v", res.Usage)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user