Skip to content
Open
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
14 changes: 14 additions & 0 deletions internal/server/api/chat.go
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,8 @@ func (handlers *ChatCompletionHandlers) ChatCompletionWithRequest(c *gin.Context
return
}

writeForwardResponseHeaders(c, result)

if result.ChatCompletion != nil {
resp := result.ChatCompletion

Expand Down Expand Up @@ -115,6 +117,18 @@ func (handlers *ChatCompletionHandlers) ChatCompletionWithRequest(c *gin.Context
}
}

func writeForwardResponseHeaders(c *gin.Context, result orchestrator.ChatCompletionResult) {
var headers http.Header
if result.ChatCompletion != nil {
headers = result.ChatCompletion.Headers
} else {
headers = httpclient.GetResponseHeaders(result.ChatCompletionStream)
}

// c.Writer.Header() is non-nil, so the merge updates it in place.
_ = httpclient.MergeForwardResponseHeaders(c.Writer.Header(), headers)
}

// StreamErrorFormatter formats a stream error into a JSON-serializable object for SSE error events.
type StreamErrorFormatter func(ctx context.Context, err error) any

Expand Down
30 changes: 30 additions & 0 deletions internal/server/api/chat_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -135,6 +135,36 @@ func TestWriteSSEStream_Success(t *testing.T) {
assert.Contains(t, body, `[DONE]`)
}

func TestWriteForwardResponseHeaders(t *testing.T) {
headers := http.Header{
httpclient.ReasoningIncludedHeader: []string{"true"},
"Set-Cookie": []string{"secret=1"},
}
tests := map[string]orchestrator.ChatCompletionResult{
"non-streaming": {
ChatCompletion: &httpclient.Response{Headers: headers},
},
"streaming": {
ChatCompletionStream: httpclient.WithResponseHeaders(
streams.SliceStream([]*httpclient.StreamEvent{}),
headers,
),
},
}

for name, result := range tests {
t.Run(name, func(t *testing.T) {
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)

writeForwardResponseHeaders(c, result)

require.Equal(t, "true", w.Header().Get(httpclient.ReasoningIncludedHeader))
require.Empty(t, w.Header().Get("Set-Cookie"))
})
}
}

func TestWriteSSEStream_ErrorFormatsAsJSON(t *testing.T) {
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
Expand Down
2 changes: 1 addition & 1 deletion llm/httpclient/client.go
Original file line number Diff line number Diff line change
Expand Up @@ -362,7 +362,7 @@ func (hc *HttpClient) DoStream(ctx context.Context, request *Request) (streams.S

stream := decoderFactory(ctx, rawResp.Body)

return stream, nil
return WithResponseHeaders(stream, MergeForwardResponseHeaders(nil, rawResp.Header)), nil
}

// BuildHttpRequest builds an HTTP request from Request.
Expand Down
12 changes: 7 additions & 5 deletions llm/httpclient/client_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -166,7 +166,7 @@ func TestHttpClientImpl_DoStream(t *testing.T) {
serverResponse func(w http.ResponseWriter, r *http.Request)
wantErr bool
wantErrContains string
validate func(stream any) bool
validate func(stream StreamDecoder) bool
}{
{
name: "successful streaming request",
Expand All @@ -189,6 +189,8 @@ func TestHttpClientImpl_DoStream(t *testing.T) {
w.Header().Set("Content-Type", "text/event-stream")
w.Header().Set("Cache-Control", "no-cache")
w.Header().Set("Connection", "keep-alive")
w.Header().Set(ReasoningIncludedHeader, "true")
w.Header().Set("Set-Cookie", "secret=1")
w.WriteHeader(http.StatusOK)

// Write SSE events
Expand All @@ -211,9 +213,9 @@ func TestHttpClientImpl_DoStream(t *testing.T) {
}
},
wantErr: false,
validate: func(stream any) bool {
// This is a basic validation - in a real test we'd iterate through the stream
return stream != nil
validate: func(stream StreamDecoder) bool {
headers := GetResponseHeaders(stream)
return headers.Get(ReasoningIncludedHeader) == "true" && headers.Get("Set-Cookie") == ""
},
},
{
Expand All @@ -230,7 +232,7 @@ func TestHttpClientImpl_DoStream(t *testing.T) {
w.Write([]byte(`{"error": "unauthorized"}`))
},
wantErr: true,
validate: func(stream any) bool {
validate: func(stream StreamDecoder) bool {
return stream == nil
},
},
Expand Down
85 changes: 85 additions & 0 deletions llm/httpclient/response_headers.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,85 @@
package httpclient

import (
"net/http"
"strings"

"github.com/looplj/axonhub/llm/streams"
)

// ReasoningIncludedHeader tells Codex that upstream usage already includes
// previously generated reasoning tokens.
const ReasoningIncludedHeader = "X-Reasoning-Included"

type responseHeadersProvider interface {
responseHeaders() http.Header
}

type responseHeadersStream struct {
streams.Stream[*StreamEvent]
headers http.Header
}

func (s *responseHeadersStream) responseHeaders() http.Header {
return s.headers
}

// WithResponseHeaders attaches HTTP response metadata to a stream without
// widening the generic streams.Stream interface.
func WithResponseHeaders(stream streams.Stream[*StreamEvent], headers http.Header) streams.Stream[*StreamEvent] {
if stream == nil || len(headers) == 0 {
return stream
}

return &responseHeadersStream{
Stream: stream,
headers: headers.Clone(),
}
}

// GetResponseHeaders returns a copy of response metadata attached to a stream.
func GetResponseHeaders(stream streams.Stream[*StreamEvent]) http.Header {
provider, ok := stream.(responseHeadersProvider)
if !ok {
return nil
}

return provider.responseHeaders().Clone()
}

// MergeForwardResponseHeaders copies the small, explicit set of upstream
// headers that are safe and meaningful at AxonHub's client boundary.
func MergeForwardResponseHeaders(dst, src http.Header) http.Header {
forward := hasOnlyTrueHeaderValues(src, ReasoningIncludedHeader)
if dst != nil {
dst.Del(ReasoningIncludedHeader)
}
if !forward {
return dst
}
if dst == nil {
dst = make(http.Header)
}

dst[ReasoningIncludedHeader] = []string{"true"}

return dst
}

func hasOnlyTrueHeaderValues(headers http.Header, name string) bool {
found := false
for key, values := range headers {
if !strings.EqualFold(key, name) {
continue
}

for _, value := range values {
found = true
if !strings.EqualFold(strings.TrimSpace(value), "true") {
return false
}
}
}

return found
}
48 changes: 48 additions & 0 deletions llm/httpclient/response_headers_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,48 @@
package httpclient

import (
"net/http"
"testing"

"github.com/stretchr/testify/require"

"github.com/looplj/axonhub/llm/streams"
)

func TestMergeForwardResponseHeaders(t *testing.T) {
tests := []struct {
name string
src http.Header
want string
}{
{name: "true", src: http.Header{"x-reasoning-included": []string{" TRUE "}}, want: "true"},
{name: "false", src: http.Header{ReasoningIncludedHeader: []string{"false"}}},
{name: "missing", src: http.Header{"Set-Cookie": []string{"secret=1"}}},
{name: "conflicting duplicates", src: http.Header{
ReasoningIncludedHeader: []string{"true"},
"x-reasoning-included": []string{"false"},
}},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
dst := http.Header{ReasoningIncludedHeader: []string{"stale"}}
got := MergeForwardResponseHeaders(dst, tt.src)

require.Equal(t, tt.want, got.Get(ReasoningIncludedHeader))
require.Empty(t, got.Get("Set-Cookie"))
})
}
}

func TestResponseHeadersStreamCopiesHeaders(t *testing.T) {
headers := http.Header{ReasoningIncludedHeader: []string{"true"}}
stream := WithResponseHeaders(streams.SliceStream([]*StreamEvent{}), headers)
headers.Set(ReasoningIncludedHeader, "false")

got := GetResponseHeaders(stream)
got.Set(ReasoningIncludedHeader, "false")

require.Equal(t, "true", GetResponseHeaders(stream).Get(ReasoningIncludedHeader))
require.Nil(t, GetResponseHeaders(streams.SliceStream([]*StreamEvent{})))
}
6 changes: 5 additions & 1 deletion llm/pipeline/integration_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -92,7 +92,9 @@ func TestPipeline_OpenAI_to_OpenAI(t *testing.T) {
return &httpclient.Response{
StatusCode: http.StatusOK,
Headers: http.Header{
"Content-Type": []string{"application/json"},
"Content-Type": []string{"application/json"},
httpclient.ReasoningIncludedHeader: []string{"true"},
"Set-Cookie": []string{"secret=1"},
},
Body: responseBody,
}, nil
Expand Down Expand Up @@ -136,6 +138,8 @@ func TestPipeline_OpenAI_to_OpenAI(t *testing.T) {
// Verify response
require.Equal(t, http.StatusOK, result.Response.StatusCode)
require.Equal(t, "application/json", result.Response.Headers.Get("Content-Type"))
require.Equal(t, "true", result.Response.Headers.Get(httpclient.ReasoningIncludedHeader))
require.Empty(t, result.Response.Headers.Get("Set-Cookie"))

var finalResponse llm.Response

Expand Down
4 changes: 4 additions & 0 deletions llm/pipeline/non_streaming.go
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,7 @@ func (p *pipeline) notStream(

return nil, fmt.Errorf("failed to apply raw response middlewares: %w", err)
}
responseHeaders := httpclient.MergeForwardResponseHeaders(nil, httpResp.Headers)

llmResp, err := p.Outbound.TransformResponse(ctx, httpResp)
if err != nil {
Expand Down Expand Up @@ -74,6 +75,7 @@ func (p *pipeline) notStream(

return nil, fmt.Errorf("failed to apply inbound raw response middlewares: %w", err)
}
finalResp.Headers = httpclient.MergeForwardResponseHeaders(finalResp.Headers, responseHeaders)

return finalResp, nil
}
Expand All @@ -88,6 +90,7 @@ func (p *pipeline) autoAggregateStream(
return nil, err
}
defer inboundStream.Close()
responseHeaders := httpclient.GetResponseHeaders(inboundStream)

chunks := make([]*httpclient.StreamEvent, 0, 8)
for inboundStream.Next() {
Expand Down Expand Up @@ -132,6 +135,7 @@ func (p *pipeline) autoAggregateStream(
p.applyRawErrorResponseMiddlewares(ctx, err)
return nil, fmt.Errorf("failed to apply inbound raw response middlewares: %w", err)
}
resp.Headers = httpclient.MergeForwardResponseHeaders(resp.Headers, responseHeaders)

return resp, nil
}
11 changes: 10 additions & 1 deletion llm/pipeline/pipeline_retry_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -470,7 +470,15 @@ func TestPipeline_Process_StreamRetriesPreCommitError(t *testing.T) {
executor := &mockExecutor{
doStream: func(ctx context.Context, req *httpclient.Request) (streams.Stream[*httpclient.StreamEvent], error) {
attempts++
return streams.SliceStream([]*httpclient.StreamEvent{{Data: []byte("raw")}}), nil
reasoningIncluded := "false"
if attempts == 1 {
reasoningIncluded = "true"
}

return httpclient.WithResponseHeaders(
streams.SliceStream([]*httpclient.StreamEvent{{Data: []byte("raw")}}),
http.Header{httpclient.ReasoningIncludedHeader: []string{reasoningIncluded}},
), nil
},
}

Expand Down Expand Up @@ -513,6 +521,7 @@ func TestPipeline_Process_StreamRetriesPreCommitError(t *testing.T) {
require.True(t, res.Stream)
require.Equal(t, 2, attempts)
require.Equal(t, 1, prepareCalls)
require.Empty(t, httpclient.GetResponseHeaders(res.EventStream).Get(httpclient.ReasoningIncludedHeader))

events, err := streams.All(res.EventStream)
require.NoError(t, err)
Expand Down
3 changes: 2 additions & 1 deletion llm/pipeline/stream.go
Original file line number Diff line number Diff line change
Expand Up @@ -321,6 +321,7 @@ func (p *pipeline) stream(

return nil, WrapUpstreamError(err)
}
responseHeaders := httpclient.MergeForwardResponseHeaders(nil, httpclient.GetResponseHeaders(outboundStream))

// Apply raw stream middlewares
rawStream := outboundStream
Expand Down Expand Up @@ -445,5 +446,5 @@ func (p *pipeline) stream(
}
}

return inboundStream, nil
return httpclient.WithResponseHeaders(inboundStream, responseHeaders), nil
}
14 changes: 12 additions & 2 deletions llm/pipeline/streaming_integration_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -72,7 +72,10 @@ func TestPipeline_Streaming_OpenAI_to_OpenAI(t *testing.T) {
require.Equal(t, "Bearer test-api-key", request.Headers.Get("Authorization"))

// Return mock stream
return streams.SliceStream(streamEvents), nil
return httpclient.WithResponseHeaders(streams.SliceStream(streamEvents), http.Header{
httpclient.ReasoningIncludedHeader: []string{"true"},
"Set-Cookie": []string{"secret=1"},
}), nil
},
}

Expand Down Expand Up @@ -128,6 +131,8 @@ func TestPipeline_Streaming_OpenAI_to_OpenAI(t *testing.T) {
require.NotNil(t, result)
require.True(t, result.Stream)
require.NotNil(t, result.EventStream)
require.Equal(t, "true", httpclient.GetResponseHeaders(result.EventStream).Get(httpclient.ReasoningIncludedHeader))
require.Empty(t, httpclient.GetResponseHeaders(result.EventStream).Get("Set-Cookie"))

// Collect all events from the stream
var collectedEvents []*httpclient.StreamEvent
Expand Down Expand Up @@ -461,7 +466,10 @@ func TestPipeline_NonStreaming_AutoAggregateUpgradedStream(t *testing.T) {
require.NoError(t, err)
require.Equal(t, true, reqBody["stream"])

return streams.SliceStream(streamEvents), nil
return httpclient.WithResponseHeaders(streams.SliceStream(streamEvents), http.Header{
httpclient.ReasoningIncludedHeader: []string{"true"},
"Set-Cookie": []string{"secret=1"},
}), nil
},
}

Expand Down Expand Up @@ -515,6 +523,8 @@ func TestPipeline_NonStreaming_AutoAggregateUpgradedStream(t *testing.T) {
require.NotNil(t, result.Response)
require.Equal(t, http.StatusOK, result.Response.StatusCode)
require.Equal(t, "application/json", result.Response.Headers.Get("Content-Type"))
require.Equal(t, "true", result.Response.Headers.Get(httpclient.ReasoningIncludedHeader))
require.Empty(t, result.Response.Headers.Get("Set-Cookie"))

var finalResponse map[string]any
err = json.Unmarshal(result.Response.Body, &finalResponse)
Expand Down
4 changes: 2 additions & 2 deletions llm/transformer/openai/codex/outbound.go
Original file line number Diff line number Diff line change
Expand Up @@ -376,9 +376,9 @@ func (e *codexExecutor) Do(ctx context.Context, request *httpclient.Request) (*h

return &httpclient.Response{
StatusCode: http.StatusOK,
Headers: http.Header{
Headers: httpclient.MergeForwardResponseHeaders(http.Header{
"Content-Type": []string{"application/json"},
},
}, httpclient.GetResponseHeaders(stream)),
Body: body,
Request: request,
}, nil
Expand Down
Loading
Loading