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
628 changes: 628 additions & 0 deletions internal/responsecache/handle_request_test.go

Large diffs are not rendered by default.

80 changes: 70 additions & 10 deletions internal/responsecache/middleware_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -200,27 +200,87 @@ func TestHashRequest_ModeChangesKey(t *testing.T) {
}
}

func TestSimpleCacheMiddleware_SkipsStreaming(t *testing.T) {
func TestHashRequest_StreamIncludeUsageChangesKey(t *testing.T) {
base := []byte(`{"model":"gpt-4","stream":true,"messages":[{"role":"user","content":"hi"}]}`)
withUsage := []byte(`{"model":"gpt-4","stream":true,"stream_options":{"include_usage":true},"messages":[{"role":"user","content":"hi"}]}`)
plan := &core.ExecutionPlan{
Mode: core.ExecutionModeTranslated,
ProviderType: "openai",
Resolution: &core.RequestModelResolution{
ResolvedSelector: core.ModelSelector{Provider: "openai", Model: "gpt-4"},
},
}

first := hashRequest("/v1/chat/completions", base, plan)
second := hashRequest("/v1/chat/completions", withUsage, plan)

if first == second {
t.Fatal("stream_options.include_usage should affect the exact cache key")
}
}

func TestSimpleCacheMiddleware_SharesCacheAcrossStreamingAndNonStreaming(t *testing.T) {
store := cache.NewMapStore()
defer store.Close()
mw := NewResponseCacheMiddlewareWithStore(store, time.Hour)
e := echo.New()
installResolvedExecutionPlan(e, "openai", "gpt-4")
e.Use(mw.Middleware())
callCount := 0
e.POST("/v1/chat/completions", func(c *echo.Context) error {
callCount++
return c.JSON(http.StatusOK, map[string]string{"n": "1"})
return c.JSON(http.StatusOK, &core.ChatResponse{
ID: "chatcmpl-shared-cache",
Object: "chat.completion",
Model: "gpt-4",
Provider: "openai",
Created: 1234567890,
Choices: []core.Choice{
{
Index: 0,
Message: core.ResponseMessage{
Role: "assistant",
Content: "shared cached response",
},
FinishReason: "stop",
},
},
Usage: core.Usage{
PromptTokens: 9,
CompletionTokens: 3,
TotalTokens: 12,
},
})
})

body := []byte(`{"model":"gpt-4","stream":true,"messages":[{"role":"user","content":"hi"}]}`)
for range 2 {
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
rec := httptest.NewRecorder()
e.ServeHTTP(rec, req)
nonStreamingBody := []byte(`{"model":"gpt-4","messages":[{"role":"user","content":"hi"}]}`)
streamingBody := []byte(`{"model":"gpt-4","stream":true,"messages":[{"role":"user","content":"hi"}]}`)

req1 := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(nonStreamingBody))
req1.Header.Set("Content-Type", "application/json")
rec1 := httptest.NewRecorder()
e.ServeHTTP(rec1, req1)
if rec1.Header().Get("X-Cache") != "" {
t.Fatalf("first request should miss cache, got X-Cache=%q", rec1.Header().Get("X-Cache"))
}
if callCount != 2 {
t.Fatalf("streaming requests should not be cached, handler called %d times", callCount)

mw.simple.wg.Wait()

req2 := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(streamingBody))
req2.Header.Set("Content-Type", "application/json")
rec2 := httptest.NewRecorder()
e.ServeHTTP(rec2, req2)
if got := rec2.Header().Get("X-Cache"); got != "HIT (exact)" {
t.Fatalf("streaming request should reuse cached full response, got X-Cache=%q", got)
}
if got := rec2.Header().Get("Content-Type"); got != "text/event-stream" {
t.Fatalf("streaming cache hit Content-Type = %q, want text/event-stream", got)
}
if !bytes.Contains(rec2.Body.Bytes(), []byte("shared cached response")) || !bytes.Contains(rec2.Body.Bytes(), []byte("[DONE]")) {
t.Fatalf("streaming cache hit body = %q, want synthesized SSE", rec2.Body.String())
}
if callCount != 1 {
t.Fatalf("expected streaming replay to avoid second handler call, got %d calls", callCount)
}
}

Expand Down
3 changes: 2 additions & 1 deletion internal/responsecache/responsecache.go
Original file line number Diff line number Diff line change
Expand Up @@ -94,7 +94,8 @@ func (m *ResponseCacheMiddleware) Middleware() echo.MiddlewareFunc {
// HandleRequest runs the full dual-layer cache check (exact then semantic) for a
// translated inference request that has already been guardrail-patched.
// body is the final patched request bytes; next is the real LLM call.
// Returns true if the request was served from cache.
// Streaming misses are reconstructed into full JSON before storage; streaming
// hits replay that stored JSON as synthetic SSE.
func (m *ResponseCacheMiddleware) HandleRequest(c *echo.Context, body []byte, next func() error) error {
if m == nil {
return next()
Expand Down
63 changes: 32 additions & 31 deletions internal/responsecache/semantic.go
Original file line number Diff line number Diff line change
Expand Up @@ -64,10 +64,6 @@ func (m *semanticCacheMiddleware) Handle(c *echo.Context, body []byte, next func
return next()
}

if isStreamingRequest(path, body) {
return next()
}

ctx := c.Request().Context()
plan := core.GetExecutionPlan(ctx)

Expand Down Expand Up @@ -105,20 +101,20 @@ func (m *semanticCacheMiddleware) Handle(c *echo.Context, body []byte, next func
}

if len(results) > 0 && float64(results[0].Score) >= threshold {
auditlog.EnrichEntryWithCacheType(c, CacheTypeSemantic)
c.Response().Header().Set("Content-Type", "application/json")
c.Response().Header().Set("X-Cache", "HIT (semantic)")
c.Response().WriteHeader(http.StatusOK)
_, _ = c.Response().Write(results[0].Response)
if m.hitRecorder != nil {
m.hitRecorder(c, results[0].Response, CacheTypeSemantic)
replayErr := writeCachedResponse(c, path, body, results[0].Response, CacheTypeSemantic)
if replayErr == nil {
auditlog.EnrichEntryWithCacheType(c, CacheTypeSemantic)
if m.hitRecorder != nil {
m.hitRecorder(c, results[0].Response, CacheTypeSemantic)
}
slog.Info("semantic cache hit",
"path", path,
"score", results[0].Score,
"request_id", c.Request().Header.Get("X-Request-ID"),
)
return nil
}
slog.Info("semantic cache hit",
"path", path,
"score", results[0].Score,
"request_id", c.Request().Header.Get("X-Request-ID"),
)
return nil
slog.Warn("semantic cache replay failed", "path", path, "err", replayErr)
}

capture := &responseCapture{
Expand All @@ -137,8 +133,11 @@ func (m *semanticCacheMiddleware) Handle(c *echo.Context, body []byte, next func
if core.GetFallbackUsed(c.Request().Context()) {
return nil
}

data := bytes.Clone(capture.body.Bytes())
data, ok := capture.cachedBody(path, streamResponseDefaultsFromContext(c.Request().Context()), c.Response().Header().Get("Content-Type"))
if !ok {
slog.Warn("semantic cache: failed to reconstruct cacheable response body", "path", path)
return nil
}
ttl := time.Duration(m.cfg.TTL) * time.Second
if ttl == 0 {
ttl = time.Hour
Expand Down Expand Up @@ -346,13 +345,13 @@ func extractTextFromContent(content any) string {
// (e.g. "/v1/chat/completions") and isolates entries across distinct endpoints.
func computeParamsHash(body []byte, endpointPath string, plan *core.ExecutionPlan, guardrailsHash, embedderIdentity string) string {
var req struct {
Model string `json:"model"`
Temperature *float64 `json:"temperature"`
TopP *float64 `json:"top_p"`
MaxTokens *int `json:"max_tokens"`
Tools []map[string]any `json:"tools"`
ResponseFormat any `json:"response_format"`
Stream bool `json:"stream"`
Model string `json:"model"`
Temperature *float64 `json:"temperature"`
TopP *float64 `json:"top_p"`
MaxTokens *int `json:"max_tokens"`
Tools []map[string]any `json:"tools"`
ResponseFormat any `json:"response_format"`
StreamOptions *core.StreamOptions `json:"stream_options"`
}
_ = json.Unmarshal(body, &req)

Expand Down Expand Up @@ -397,7 +396,10 @@ func computeParamsHash(body []byte, endpointPath string, plan *core.ExecutionPla
}
h.Write([]byte{0})

h.Write([]byte(strconv.FormatBool(req.Stream)))
if streamOptions := normalizeStreamOptionsForCache(req.StreamOptions); streamOptions != nil {
soJSON, _ := json.Marshal(streamOptions)
h.Write(soJSON)
}
h.Write([]byte{0})

h.Write([]byte(guardrailsHash))
Expand Down Expand Up @@ -512,9 +514,8 @@ func ShouldSkipExactCache(req *http.Request) bool {
return strings.EqualFold(req.Header.Get("X-Cache-Type"), CacheTypeSemantic)
}

// ShouldSkipAllCache reports whether caching must be bypassed for this request
// (X-Cache-Control: no-store or Cache-Control containing no-store), matching
// shouldSkipSemanticCache / shouldSkipCache semantics for the no-store directive.
// ShouldSkipAllCache reports whether caching must be bypassed for this request,
// matching the exact-cache middleware semantics for no-cache and no-store.
func ShouldSkipAllCache(req *http.Request) bool {
if strings.EqualFold(req.Header.Get("X-Cache-Control"), "no-store") {
return true
Expand All @@ -526,7 +527,7 @@ func ShouldSkipAllCache(req *http.Request) bool {
directives := strings.Split(strings.ToLower(cc), ",")
for _, d := range directives {
d = strings.TrimSpace(d)
if d == "no-store" {
if d == "no-cache" || d == "no-store" {
return true
}
}
Expand Down
92 changes: 87 additions & 5 deletions internal/responsecache/semantic_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -207,6 +207,25 @@ func TestSemanticCacheMiddleware_ParamsHashIsolation_Temperature(t *testing.T) {
}
}

func TestComputeParamsHash_StreamIncludeUsageChangesHash(t *testing.T) {
base := []byte(`{"model":"gpt-4","stream":true,"messages":[{"role":"user","content":"same prompt"}]}`)
withUsage := []byte(`{"model":"gpt-4","stream":true,"stream_options":{"include_usage":true},"messages":[{"role":"user","content":"same prompt"}]}`)
plan := &core.ExecutionPlan{
Mode: core.ExecutionModeTranslated,
ProviderType: "openai",
Resolution: &core.RequestModelResolution{
ResolvedSelector: core.ModelSelector{Provider: "openai", Model: "gpt-4"},
},
}

first := computeParamsHash(base, "/v1/chat/completions", plan, "", "")
second := computeParamsHash(withUsage, "/v1/chat/completions", plan, "", "")

if first == second {
t.Fatal("stream_options.include_usage should affect semantic params_hash")
}
}

func TestSemanticCacheMiddleware_GuardrailsHashIsolation(t *testing.T) {
m, store, emb := newTestSemanticMiddleware(0.90, 10, false)
emb.vector = []float32{1, 0, 0}
Expand Down Expand Up @@ -284,15 +303,70 @@ func TestSemanticCacheMiddleware_ExcludeSystemPrompt(t *testing.T) {
}
}

func TestSemanticCacheMiddleware_StreamingSkipped(t *testing.T) {
func TestSemanticCacheMiddleware_StreamingMissPopulatesSemanticCacheAcrossModes(t *testing.T) {
m, store, _ := newTestSemanticMiddleware(0.90, 10, false)
streamBody := []byte(`{"model":"gpt-4","stream":true,"messages":[{"role":"user","content":"semantic-stream-cache"}]}`)
jsonBody := []byte(`{"model":"gpt-4","messages":[{"role":"user","content":"semantic-stream-cache"}]}`)
e := echo.New()
handlerCalls := 0

run := func(body []byte) *httptest.ResponseRecorder {
t.Helper()
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
if err := m.Handle(c, body, func() error {
handlerCalls++
c.Response().Header().Set("Content-Type", "text/event-stream")
c.Response().WriteHeader(http.StatusOK)
_, _ = c.Response().Write([]byte("data: {\"id\":\"chatcmpl-semantic-stream\",\"object\":\"chat.completion.chunk\",\"created\":1234567890,\"model\":\"gpt-4\",\"provider\":\"openai\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"Semantic\"},\"finish_reason\":null}]}\n\n"))
_, _ = c.Response().Write([]byte("data: {\"id\":\"chatcmpl-semantic-stream\",\"object\":\"chat.completion.chunk\",\"created\":1234567890,\"model\":\"gpt-4\",\"provider\":\"openai\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" cache\"},\"finish_reason\":\"stop\"}],\"usage\":{\"prompt_tokens\":7,\"completion_tokens\":2,\"total_tokens\":9}}\n\n"))
_, _ = c.Response().Write([]byte("data: [DONE]\n\n"))
return nil
}); err != nil {
t.Fatalf("Handle error: %v", err)
}
return rec
}

rec1 := run(streamBody)
if got := rec1.Header().Get("X-Cache"); got != "" {
t.Fatalf("streaming miss should not be cache hit, got X-Cache=%q", got)
}

body := []byte(`{"model":"gpt-4","stream":true,"messages":[{"role":"user","content":"hi"}]}`)
serveSemanticRequest(t, m, body, "")
m.wg.Wait()

if store.Len() != 0 {
t.Fatal("streaming requests should be skipped by semantic cache")
if store.Len() != 1 {
t.Fatalf("expected streaming miss to populate semantic cache, got %d entries", store.Len())
}

rec2 := run(jsonBody)
if got := rec2.Header().Get("X-Cache"); got != "HIT (semantic)" {
t.Fatalf("non-streaming follow-up should hit semantic cache, got X-Cache=%q", got)
}
if got := rec2.Header().Get("Content-Type"); got != "application/json" {
t.Fatalf("non-streaming semantic hit Content-Type = %q, want application/json", got)
}
if !bytes.Contains(rec2.Body.Bytes(), []byte(`"content":"Semantic cache"`)) {
t.Fatalf("semantic cache hit body = %q, want reconstructed JSON response", rec2.Body.String())
}
if handlerCalls != 1 {
t.Fatalf("semantic hit should not call handler again, got %d calls", handlerCalls)
}

rec3 := run(streamBody)
if got := rec3.Header().Get("X-Cache"); got != "HIT (semantic)" {
t.Fatalf("streaming follow-up should hit semantic cache, got X-Cache=%q", got)
}
if got := rec3.Header().Get("Content-Type"); got != "text/event-stream" {
t.Fatalf("streaming semantic hit Content-Type = %q, want text/event-stream", got)
}
if !bytes.Contains(rec3.Body.Bytes(), []byte("Semantic cache")) || !bytes.Contains(rec3.Body.Bytes(), []byte("[DONE]")) {
t.Fatalf("streaming semantic hit body = %q, want synthesized SSE", rec3.Body.String())
}
if handlerCalls != 1 {
t.Fatalf("streaming semantic hit should not call handler again, got %d calls", handlerCalls)
}
}

Expand Down Expand Up @@ -462,6 +536,14 @@ func TestShouldSkipAllCache_CacheControlNoStore(t *testing.T) {
}
}

func TestShouldSkipAllCache_CacheControlNoCache(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil)
req.Header.Set("Cache-Control", "private, no-cache, max-age=0")
if !ShouldSkipAllCache(req) {
t.Fatal("expected ShouldSkipAllCache for Cache-Control: no-cache")
}
}

func TestSemanticCacheMiddleware_HitMarksAuditEntryCacheType(t *testing.T) {
m, _, _ := newTestSemanticMiddleware(0.90, 10, false)
body := []byte(`{"model":"gpt-4","messages":[{"role":"user","content":"semantic-cache-type"}]}`)
Expand Down
Loading
Loading