Add direct per-run session overrides

This commit is contained in:
2026-07-29 19:43:26 +00:00
parent eb8ab215e8
commit f6ee18f6b3
13 changed files with 302 additions and 27 deletions

View File

@@ -63,6 +63,7 @@ type RunRequest struct {
PromptID string
PromptVersion string
ProfileID string
SessionID string
APIKey string `json:"-" yaml:"-"`
Inputs map[string]ArtifactRef
Vars map[string]string
@@ -79,6 +80,7 @@ type RunResult struct {
PromptID string
PromptVersion string
PromptHash string
SessionID string
RenderedPromptHash string
SelectedProfileID string
SelectedBackendID string

View File

@@ -0,0 +1,19 @@
package domain
import (
"fmt"
"strings"
"unicode/utf8"
)
// NormalizeSessionID applies the shared session identifier rule.
func NormalizeSessionID(raw string) (string, error) {
normalized := strings.TrimSpace(raw)
if normalized == "" {
return "", nil
}
if length := utf8.RuneCountInString(normalized); length > SessionIDMaxLength {
return "", fmt.Errorf("session_id length %d exceeds maximum %d", length, SessionIDMaxLength)
}
return normalized, nil
}

View File

@@ -0,0 +1,57 @@
package domain
import (
"strings"
"testing"
)
func TestNormalizeSessionID(t *testing.T) {
tests := []struct {
name string
raw string
want string
wantErr bool
}{
{
name: "trims surrounding Unicode whitespace",
raw: "\u2003 session-123 \u2003",
want: "session-123",
},
{
name: "blank input is omitted",
raw: " \t\u2003 ",
want: "",
},
{
name: "maximum Unicode length is accepted",
raw: strings.Repeat("界", SessionIDMaxLength),
want: strings.Repeat("界", SessionIDMaxLength),
},
{
name: "one Unicode code point over maximum is rejected",
raw: strings.Repeat("界", SessionIDMaxLength+1),
wantErr: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := NormalizeSessionID(tt.raw)
if tt.wantErr {
if err == nil {
t.Fatal("expected normalization error")
}
if !strings.Contains(err.Error(), "exceeds maximum") {
t.Fatalf("expected useful length diagnostic, got %v", err)
}
return
}
if err != nil {
t.Fatalf("normalize session id: %v", err)
}
if got != tt.want {
t.Fatalf("normalized session id = %q, want %q", got, tt.want)
}
})
}
}

View File

@@ -12,7 +12,6 @@ import (
"os"
"strings"
"time"
"unicode/utf8"
"gitea.maximumdirect.net/eric/promptkit/internal/defaults"
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
@@ -177,12 +176,11 @@ func openAIChatRequestFromGenerateRequest(req domain.GenerateRequest, defaultMod
wireReq := openAIChatRequest{
Model: model,
}
if sessionID := strings.TrimSpace(req.Prompt.SessionID); sessionID != "" {
if n := utf8.RuneCountInString(sessionID); n > domain.SessionIDMaxLength {
return openAIChatRequest{}, fmt.Errorf("session_id length %d exceeds maximum %d", n, domain.SessionIDMaxLength)
}
wireReq.SessionID = sessionID
sessionID, err := domain.NormalizeSessionID(req.Prompt.SessionID)
if err != nil {
return openAIChatRequest{}, err
}
wireReq.SessionID = sessionID
wireReq.Messages = make([]openAIChatRequestMessage, 0, len(req.Prompt.Messages))
for _, msg := range req.Prompt.Messages {

View File

@@ -5,10 +5,9 @@ import (
"context"
"errors"
"fmt"
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
"strings"
"text/template"
"unicode/utf8"
"gitea.maximumdirect.net/eric/promptkit/internal/domain"
)
var (
@@ -95,10 +94,6 @@ func (r *goRenderer) Render(ctx context.Context, definition *domain.PromptDefini
}
func renderSessionID(raw string, funcs template.FuncMap, vars map[string]string) (string, error) {
if strings.TrimSpace(raw) == "" {
return "", nil
}
tmpl, err := template.New("session_id").Funcs(funcs).Option("missingkey=error").Parse(raw)
if err != nil {
return "", fmt.Errorf("%w: session_id: %v", ErrInvalidTemplate, err)
@@ -109,9 +104,9 @@ func renderSessionID(raw string, funcs template.FuncMap, vars map[string]string)
return "", fmt.Errorf("%w: session_id: %w", ErrRenderFailure, err)
}
sessionID := strings.TrimSpace(buf.String())
if n := utf8.RuneCountInString(sessionID); n > domain.SessionIDMaxLength {
return "", fmt.Errorf("%w: session_id length %d exceeds maximum %d", ErrRenderFailure, n, domain.SessionIDMaxLength)
sessionID, err := domain.NormalizeSessionID(buf.String())
if err != nil {
return "", fmt.Errorf("%w: session_id: %v", ErrRenderFailure, err)
}
return sessionID, nil
}

View File

@@ -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{

View File

@@ -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)
}