Implement support for OpenRouter sticky routing via a session_id variable

This commit is contained in:
2026-07-02 20:08:44 -05:00
parent 4d4bb7a121
commit 63fb8fc132
19 changed files with 363 additions and 13 deletions

View File

@@ -12,6 +12,7 @@ import (
"os"
"strings"
"time"
"unicode/utf8"
"gitea.maximumdirect.net/eric/scriptorium/internal/defaults"
"gitea.maximumdirect.net/eric/scriptorium/internal/domain"
@@ -169,6 +170,12 @@ 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
}
wireReq.Messages = make([]openAIChatRequestMessage, 0, len(req.Prompt.Messages))
for _, msg := range req.Prompt.Messages {
@@ -200,6 +207,7 @@ func openAIChatRequestFromGenerateRequest(req domain.GenerateRequest, defaultMod
type openAIChatRequest struct {
Model string `json:"model"`
SessionID string `json:"session_id,omitempty"`
Messages []openAIChatRequestMessage `json:"messages"`
Temperature *float64 `json:"temperature,omitempty"`
MaxTokens *int `json:"max_tokens,omitempty"`

View File

@@ -181,7 +181,7 @@ func TestOpenAICompatibleClientSerializesCacheControlledMessageAsContentBlock(t
t.Fatalf("expected no error, got %v", err)
}
for _, forbidden := range []string{"cache_control", "session_id", "extra_params"} {
for _, forbidden := range []string{"cache_control", "extra_params"} {
if _, exists := observedBody[forbidden]; exists {
t.Fatalf("expected top-level %s to be omitted, got %#v", forbidden, observedBody[forbidden])
}
@@ -262,6 +262,96 @@ func TestOpenAICompatibleClientOmitsEmptyCacheControlTTL(t *testing.T) {
}
}
func TestOpenAICompatibleClientSerializesSessionID(t *testing.T) {
var observedBody map[string]any
var observedSessionHeader string
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
observedSessionHeader = r.Header.Get("x-session-id")
defer r.Body.Close()
if err := json.NewDecoder(r.Body).Decode(&observedBody); err != nil {
t.Fatalf("failed to decode request body: %v", err)
}
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`))
}))
defer ts.Close()
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1"})
if err != nil {
t.Fatal(err)
}
_, err = client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{
SessionID: " session-123 ",
Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}},
},
Target: domain.ExecutionTarget{Model: "model"},
})
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
if observedBody["session_id"] != "session-123" {
t.Fatalf("expected top-level session_id, got %#v", observedBody["session_id"])
}
if observedSessionHeader != "" {
t.Fatalf("did not expect x-session-id header, got %q", observedSessionHeader)
}
}
func TestOpenAICompatibleClientOmitsEmptySessionID(t *testing.T) {
var observedBody map[string]any
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
defer r.Body.Close()
if err := json.NewDecoder(r.Body).Decode(&observedBody); err != nil {
t.Fatalf("failed to decode request body: %v", err)
}
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"ok"}}]}`))
}))
defer ts.Close()
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{BaseURL: ts.URL + "/v1"})
if err != nil {
t.Fatal(err)
}
_, err = client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{
SessionID: " ",
Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}},
},
Target: domain.ExecutionTarget{Model: "model"},
})
if err != nil {
t.Fatalf("expected no error, got %v", err)
}
if _, exists := observedBody["session_id"]; exists {
t.Fatalf("expected empty session_id to be omitted, got %#v", observedBody["session_id"])
}
}
func TestOpenAICompatibleClientRejectsTooLongSessionID(t *testing.T) {
client, err := NewOpenAICompatibleClient(OpenAICompatibleConfig{
BaseURL: "http://example.com/v1",
Model: "model",
})
if err != nil {
t.Fatal(err)
}
_, err = client.Generate(context.Background(), domain.GenerateRequest{
Prompt: domain.RenderedPrompt{
SessionID: strings.Repeat("x", domain.SessionIDMaxLength+1),
Messages: []domain.RenderedMessage{{Role: "user", Content: "hi"}},
},
})
if err == nil {
t.Fatal("expected invalid request error")
}
if !errors.Is(err, ErrInvalidRequest) {
t.Fatalf("expected ErrInvalidRequest, got %v", err)
}
}
func TestOpenAICompatibleClientParsesCacheUsage(t *testing.T) {
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = w.Write([]byte(`{