diff --git a/ext/ext.go b/ext/ext.go index 33a9f906..a5172f02 100644 --- a/ext/ext.go +++ b/ext/ext.go @@ -41,6 +41,10 @@ type Input struct { UserPath string // RequestID is the request correlation ID (X-Request-ID). RequestID string + // SessionID is the detected client session, when present, already scoped + // by the effective user path. Session detection runs before rewriters, so + // a rewriter can keep its decisions stable across a conversation. + SessionID string } // Result carries a rewritten body and response-header annotations. diff --git a/internal/server/request_rewrite.go b/internal/server/request_rewrite.go index 8ba39d30..3543b0be 100644 --- a/internal/server/request_rewrite.go +++ b/internal/server/request_rewrite.go @@ -42,6 +42,7 @@ func RequestRewriteMiddleware(rewriters []ext.RequestRewriter, auditLogger audit Header: redactCredentialHeaders(c.Request().Header), UserPath: core.UserPathFromContext(c.Request().Context()), RequestID: core.GetRequestID(c.Request().Context()), + SessionID: core.SessionIDFromContext(c.Request().Context()), } changed := false diff --git a/internal/server/request_rewrite_test.go b/internal/server/request_rewrite_test.go index 8a8dc769..d14fb16f 100644 --- a/internal/server/request_rewrite_test.go +++ b/internal/server/request_rewrite_test.go @@ -119,6 +119,59 @@ func TestRequestRewriteMiddlewareRewritesChatCompletions(t *testing.T) { } } +func TestRequestRewriteMiddlewareExposesSessionID(t *testing.T) { + tests := []struct { + name string + session string // stamped into the context before rewriters; "" = not detected + want string + }{ + {name: "detected session propagates", session: "sess-42", want: "sess-42"}, + {name: "no detected session yields empty", session: "", want: ""}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + provider := newRewriteTestProvider() + var seenSession string + capturing := &stubRewriter{ + name: "capture-session", + rewrite: func(in ext.Input) (*ext.Result, error) { + seenSession = in.SessionID + return nil, nil + }, + } + // Session detection runs before rewriters and stamps the request + // context; ExtraMiddleware runs even earlier, so it stands in for + // the detector here. + stampSession := func(next echo.HandlerFunc) echo.HandlerFunc { + return func(c *echo.Context) error { + if tt.session != "" { + req := c.Request() + c.SetRequest(req.WithContext(core.WithSessionID(req.Context(), tt.session))) + } + return next(c) + } + } + srv := New(provider, &Config{ + RequestRewriters: []ext.RequestRewriter{capturing}, + ExtraMiddleware: []echo.MiddlewareFunc{stampSession}, + }) + + rec := postJSON(t, srv, "/v1/chat/completions", + `{"model":"gpt-4o-mini","messages":[{"role":"user","content":"hi"}]}`) + if rec.Code != http.StatusOK { + t.Fatalf("expected 200, got %d (%s)", rec.Code, rec.Body.String()) + } + // The empty-session expectation must not pass vacuously. + if capturing.calls != 1 { + t.Fatalf("rewriter called %d times, want 1", capturing.calls) + } + if seenSession != tt.want { + t.Errorf("rewriter saw SessionID %q, want %q", seenSession, tt.want) + } + }) + } +} + func TestRequestRewriteMiddlewareRewritesMessages(t *testing.T) { provider := newRewriteTestProvider() srv := New(provider, &Config{