Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 39 additions & 0 deletions ext/ext.go
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,45 @@ type RequestRewriter interface {
Rewrite(ctx context.Context, in Input) (*Result, error)
}

// ResponseFeedbackObserver receives content-free feedback after a rewritten
// request completes successfully at a provider. For ordinary responses the
// callback runs after the provider returns; for streaming responses it runs
// after the stream closes. Failed requests do not produce feedback.
//
// It is an optional companion to RequestRewriter: core detects implementations
// structurally. usageObserved distinguishes a provider-confirmed zero from a
// provider or stream that returned no usage breakdown. Implementations can be
// called concurrently, must therefore be concurrency-safe, and must return
// promptly so they do not delay completion of the client request.
//
// The flat signature intentionally uses only long-standing extension types so
// extensions can implement the hook while supporting older core releases;
// older cores simply never call it.
type ResponseFeedbackObserver interface {
ObserveResponse(
ctx context.Context,
requestID string,
endpoint Endpoint,
sessionID string,
model string,
providerType string,
providerName string,
inputTokens int,
cachedInputTokens int,
cacheWriteInputTokens int,
usageObserved bool,
)
}

// ResponseFeedbackFilter lets an observer control registration per request.
// Core calls it after Rewrite and registers the observer only when it returns
// true. Observers without this optional interface receive feedback for every
// successful rewritten request. A panic is isolated and treated as false so
// this optional hook cannot abort inference.
type ResponseFeedbackFilter interface {
WantsResponseFeedback(in Input, result *Result) bool
}

Comment thread
coderabbitai[bot] marked this conversation as resolved.
// SettingOption is one allowed value for a dashboard-editable extension
// setting. Label and Description are safe to expose in the admin UI.
type SettingOption struct {
Expand Down
29 changes: 29 additions & 0 deletions internal/server/request_rewrite.go
Original file line number Diff line number Diff line change
Expand Up @@ -47,11 +47,21 @@ func RequestRewriteMiddleware(rewriters []ext.RequestRewriter, auditLogger audit

changed := false
tokensSaved := 0
feedbackObservers := make([]ext.ResponseFeedbackObserver, 0, len(rewriters))
for _, rw := range rewriters {
res, rwErr := rw.Rewrite(c.Request().Context(), in)
if rwErr != nil {
return handleError(c, rewriterGatewayError(rw.Name(), rwErr))
}
if observer, ok := rw.(ext.ResponseFeedbackObserver); ok {
wantsFeedback := true
if filter, filtered := rw.(ext.ResponseFeedbackFilter); filtered {
wantsFeedback = safelyWantsResponseFeedback(rw.Name(), filter, in, res)
}
if wantsFeedback {
feedbackObservers = append(feedbackObservers, observer)
}
}
if res != nil {
applyRewriteResponseHeaders(c, res.ResponseHeader)
}
Expand All @@ -78,11 +88,30 @@ func RequestRewriteMiddleware(rewriters []ext.RequestRewriter, auditLogger audit
c.SetRequest(req.WithContext(core.WithRewriteTokensSaved(req.Context(), tokensSaved)))
}
}
if len(feedbackObservers) > 0 {
setResponseFeedbackObservers(c, feedbackObservers)
}
return next(c)
}
}
}

// safelyWantsResponseFeedback isolates an optional extension hook from the
// inference path. A broken filter declines feedback for this request but must
// never prevent the provider call itself.
func safelyWantsResponseFeedback(name string, filter ext.ResponseFeedbackFilter, in ext.Input, res *ext.Result) (wants bool) {
wants = true
defer func() {
if recovered := recover(); recovered != nil {
wants = false
slog.Warn("response feedback filter panicked; feedback disabled for request",
"rewriter", name,
)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
}
}()
return filter.WantsResponseFeedback(in, res)
}

// redactCredentialHeaders clones the request headers with credential values
// (Authorization, cookies, API keys, ...) masked. Rewriters run post-auth and
// get UserPath for identity, so they never need raw credentials — and this
Expand Down
97 changes: 97 additions & 0 deletions internal/server/request_rewrite_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@ import (
"github.com/enterpilot/gomodel/ext"
"github.com/enterpilot/gomodel/internal/auditlog"
"github.com/enterpilot/gomodel/internal/core"
"github.com/enterpilot/gomodel/internal/session"
)

type stubRewriter struct {
Expand All @@ -23,6 +24,24 @@ type stubRewriter struct {
rewrite func(in ext.Input) (*ext.Result, error)
}

type feedbackRewriter struct {
stubRewriter
feedbackCaptureObserver
}

type filteredFeedbackRewriter struct {
feedbackRewriter
want bool
panicFilter bool
}

func (r *filteredFeedbackRewriter) WantsResponseFeedback(ext.Input, *ext.Result) bool {
if r.panicFilter {
panic("feedback filter failed")
}
return r.want
}

func (r *stubRewriter) Name() string { return r.name }

func (r *stubRewriter) Rewrite(_ context.Context, in ext.Input) (*ext.Result, error) {
Expand Down Expand Up @@ -119,6 +138,84 @@ func TestRequestRewriteMiddlewareRewritesChatCompletions(t *testing.T) {
}
}

func TestRequestRewriteMiddlewareDeliversProviderFeedback(t *testing.T) {
provider := newRewriteTestProvider()
provider.response.Usage = core.Usage{
PromptTokens: 2048,
PromptTokensDetails: &core.PromptTokensDetails{CachedTokens: 1536},
}
rewriter := &feedbackRewriter{stubRewriter: stubRewriter{name: "feedback"}}
srv := New(provider, &Config{
RequestRewriters: []ext.RequestRewriter{rewriter},
SessionDetector: session.NewDetector(session.BuiltinRules(), false),
})

req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions",
strings.NewReader(`{"model":"gpt-4o-mini","messages":[{"role":"user","content":"hello"}]}`))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-Session-Id", "feedback-session")
req.Header.Set("X-Request-Id", "feedback-request")
rec := httptest.NewRecorder()
srv.ServeHTTP(rec, req)

if rec.Code != http.StatusOK {
t.Fatalf("status = %d (%s)", rec.Code, rec.Body.String())
}
if len(rewriter.feedback) != 1 {
t.Fatalf("feedback count = %d, want 1", len(rewriter.feedback))
}
got := rewriter.feedback[0]
if got.requestID != "feedback-request" || got.sessionID != "feedback-session" || got.cacheRead != 1536 || !got.usageObserved {
t.Fatalf("feedback = %+v", got)
}
}

func TestRequestRewriteMiddlewareHonorsResponseFeedbackFilter(t *testing.T) {
for _, want := range []bool{false, true} {
rewriter := &filteredFeedbackRewriter{
feedbackRewriter: feedbackRewriter{stubRewriter: stubRewriter{name: "filtered"}},
want: want,
}
var attached bool
next := func(c *echo.Context) error {
attached = hasResponseFeedbackObservers(c)
return c.NoContent(http.StatusOK)
}
e := echo.New()
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions",
strings.NewReader(`{"model":"gpt-4o-mini","messages":[]}`))
c := e.NewContext(req, httptest.NewRecorder())
if err := RequestRewriteMiddleware([]ext.RequestRewriter{rewriter}, nil)(next)(c); err != nil {
t.Fatalf("want=%v: middleware: %v", want, err)
}
if attached != want {
t.Fatalf("want=%v: attached=%v", want, attached)
}
}
}

func TestRequestRewriteMiddlewareIsolatesResponseFeedbackFilterPanic(t *testing.T) {
provider := newRewriteTestProvider()
rewriter := &filteredFeedbackRewriter{
feedbackRewriter: feedbackRewriter{stubRewriter: stubRewriter{name: "panicking-filter"}},
panicFilter: true,
}
srv := New(provider, &Config{RequestRewriters: []ext.RequestRewriter{rewriter}})

rec := postJSON(t, srv, "/v1/chat/completions",
`{"model":"gpt-4o-mini","messages":[{"role":"user","content":"hello"}]}`)

if rec.Code != http.StatusOK {
t.Fatalf("status = %d (%s), want 200", rec.Code, rec.Body.String())
}
if provider.capturedChatReq == nil {
t.Fatal("provider was not called after feedback filter panic")
}
if len(rewriter.feedback) != 0 {
t.Fatalf("feedback count = %d, want 0 after filter panic", len(rewriter.feedback))
}
}

func TestRequestRewriteMiddlewareExposesSessionID(t *testing.T) {
tests := []struct {
name string
Expand Down
Loading