package openai import ( "bytes" "context" "encoding/json" "errors" "fmt" "io" "slices" "strings" "sync" "testing" "time" "iop/packages/go/config" "iop/packages/go/streamgate" iop "iop/proto/gen/iop" ) type openAIResumeTestEventSource struct{} func (openAIResumeTestEventSource) NextEvent(context.Context) (streamgate.NormalizedEvent, error) { return streamgate.NormalizedEvent{}, io.EOF } type openAIResumeTrace struct { steps []string } func (t *openAIResumeTrace) record(step string) { t.steps = append(t.steps, step) } func (t *openAIResumeTrace) snapshot() []string { return append([]string(nil), t.steps...) } type openAIResumeTestController struct { calls int trace *openAIResumeTrace } func (c *openAIResumeTestController) AbortAttempt(context.Context) error { c.calls++ if c.trace != nil { c.trace.record("abort") } return nil } type openAIResumeTracingRebuilder struct { inner streamgate.RequestRebuilder trace *openAIResumeTrace } func (r openAIResumeTracingRebuilder) RebuildRequest(ctx context.Context, snapshot streamgate.RecoveryRequestSnapshotRef, plan streamgate.RecoveryPlan) (streamgate.RebuiltRequestDraft, error) { r.trace.record("rebuild") return r.inner.RebuildRequest(ctx, snapshot, plan) } type openAIResumeTracingDispatcher struct { inner streamgate.AttemptDispatcher trace *openAIResumeTrace } func (d openAIResumeTracingDispatcher) DispatchAttempt(ctx context.Context, request streamgate.RebuiltRequest) (streamgate.AttemptBinding, error) { d.trace.record("dispatch") return d.inner.DispatchAttempt(ctx, request) } type openAIResumeNoDispatch struct{ calls int } func (d *openAIResumeNoDispatch) DispatchAttempt(context.Context, streamgate.RebuiltRequest) (streamgate.AttemptBinding, error) { d.calls++ return streamgate.AttemptBinding{}, errors.New("context-overflow recovery must not dispatch") } type openAIResumeRecordingPreparer struct{ calls int } func (p *openAIResumeRecordingPreparer) PrepareRecoveryPlan(context.Context, streamgate.RecoveryPlan, streamgate.RecoveryPreparationSnapshot) (streamgate.RecoveryDirective, error) { p.calls++ return streamgate.RecoveryDirective{}, errors.New("resume continuation must not invoke a preparer") } func mustOpenAIRecoveryPlan(t *testing.T, id string, strategy streamgate.RecoveryStrategy, directive streamgate.RecoveryDirective) streamgate.RecoveryPlan { t.Helper() intent, err := streamgate.NewRecoveryIntent(strategy, directive, "openai.recovery", 10) if err != nil { t.Fatalf("NewRecoveryIntent: %v", err) } contributor, err := streamgate.NewRecoveryContributor("openai.edge", "filter.test", "rule.test") if err != nil { t.Fatalf("NewRecoveryContributor: %v", err) } policy, err := streamgate.NewRecoveryPolicySnapshot(3, map[streamgate.RecoveryStrategy]int{ streamgate.RecoveryStrategyExactReplay: 3, streamgate.RecoveryStrategyContinuationRepair: 3, streamgate.RecoveryStrategySchemaRepair: 3, }) if err != nil { t.Fatalf("NewRecoveryPolicySnapshot: %v", err) } usage, err := streamgate.NewRecoveryUsageSnapshot(0, nil) if err != nil { t.Fatalf("NewRecoveryUsageSnapshot: %v", err) } commit := streamgate.CommitStateTransportUncommitted if strategy == streamgate.RecoveryStrategyContinuationRepair { commit = streamgate.CommitStateStreamOpen } plan, err := streamgate.NewRecoveryPlan(streamgate.RecoveryEligibilityInput{ PlanID: id, IdempotencyKey: id + ":idempotency", Intent: intent, Contributors: []streamgate.RecoveryContributor{contributor}, CommitState: commit, Policy: policy, Usage: usage, }) if err != nil { t.Fatalf("NewRecoveryPlan: %v", err) } return plan } func newOpenAIRebuilderFixture(t *testing.T, endpoint string, body []byte, maxBytes int64) (*openAIIngressSnapshot, *openAIRequestRebuilder, streamgate.RecoveryRequestSnapshotRef) { t.Helper() ingress, err := buildOpenAIIngressSnapshot(maxBytes, body, json.RawMessage(body)) if err != nil { t.Fatalf("buildOpenAIIngressSnapshot: %v", err) } t.Cleanup(ingress.Close) rebuilder, err := newOpenAIRequestRebuilder(ingress, endpoint, nil, 0) if err != nil { t.Fatalf("newOpenAIRequestRebuilder: %v", err) } t.Cleanup(rebuilder.Close) ref, err := ingress.recoveryRef() if err != nil { t.Fatalf("recoveryRef: %v", err) } return ingress, rebuilder, ref } func newOpenAIResumeRebuilderFixture(t *testing.T, endpoint string, body []byte, contextWindowTokens int) (*openAIIngressSnapshot, *openAIRequestRebuilder, *openAIRecoverySourceStore, streamgate.RecoveryRequestSnapshotRef) { t.Helper() ingress, err := buildOpenAIIngressSnapshot(8192, body, json.RawMessage(body)) if err != nil { t.Fatalf("buildOpenAIIngressSnapshot: %v", err) } t.Cleanup(ingress.Close) source := newOpenAIRecoverySourceStore(ingress) rebuilder, err := newOpenAIRequestRebuilder(ingress, endpoint, source, contextWindowTokens) if err != nil { t.Fatalf("newOpenAIRequestRebuilder: %v", err) } t.Cleanup(rebuilder.Close) ref, err := ingress.recoveryRef() if err != nil { t.Fatalf("recoveryRef: %v", err) } return ingress, rebuilder, source, ref } func TestOpenAIRequestRebuilderBuildsChatRepeatResume(t *testing.T) { const callerSentinel = "CALLER_CHAT_HISTORY_MUST_NOT_BE_COPIED" const content = "MODEL_CONTENT_PREFIX__REPEATED_TAIL" const reasoning = "MODEL_REASONING_SENTINEL" body := []byte(`{"model":"served-chat","stream":true,"messages":[{"role":"user","content":"` + callerSentinel + `"}],"tools":[{"type":"function"}]}`) _, rebuilder, source, ref := newOpenAIResumeRebuilderFixture(t, openAIRebuildEndpointChat, body, 4096) source.recordText(content) source.recordReasoning(reasoning) directive, err := streamgate.NewRecoveryDirectiveContinuation(len("MODEL_CONTENT_PREFIX__"), source.snapshotRef()) if err != nil { t.Fatalf("NewRecoveryDirectiveContinuation: %v", err) } draft, err := rebuilder.RebuildRequest(context.Background(), ref, mustOpenAIRecoveryPlan(t, "plan.chat.resume", streamgate.RecoveryStrategyContinuationRepair, directive)) if err != nil { t.Fatalf("RebuildRequest: %v", err) } lease, err := rebuilder.RebuiltStore().take(draft.RequestRef()) if err != nil { t.Fatalf("take: %v", err) } defer lease.release() got, err := lease.body() if err != nil { t.Fatalf("body: %v", err) } if bytes.Contains(got, []byte(callerSentinel)) || bytes.Contains(got, []byte(`"tools"`)) { t.Fatalf("chat resume copied caller history: %s", got) } var rebuilt struct { Model string `json:"model"` Stream bool `json:"stream"` Messages []struct { Role string `json:"role"` Content string `json:"content"` ReasoningContent string `json:"reasoning_content"` } `json:"messages"` } if err := json.Unmarshal(got, &rebuilt); err != nil { t.Fatalf("decode resume body: %v", err) } if rebuilt.Model != "served-chat" || !rebuilt.Stream || len(rebuilt.Messages) != 2 { t.Fatalf("chat resume envelope = %#v", rebuilt) } if got := rebuilt.Messages[0]; got.Role != "assistant" || got.Content != "MODEL_CONTENT_PREFIX__" || got.ReasoningContent != reasoning { t.Fatalf("chat assistant provenance = %#v", got) } if got := rebuilt.Messages[1]; got.Role != "user" || got.Content != openAIRepeatResumeDirective { t.Fatalf("chat resume directive = %#v", got) } if _, err := rebuilder.RebuildRequest(context.Background(), ref, mustOpenAIRecoveryPlan(t, "plan.chat.resume.second", streamgate.RecoveryStrategyContinuationRepair, directive)); !errors.Is(err, errOpenAIRecoverySourceUnavailable) { t.Fatalf("second resume error = %v, want one-shot source error", err) } } func TestOpenAIRequestRebuilderBuildsResponsesRepeatResume(t *testing.T) { const callerInput = "CALLER_RESPONSES_INPUT_MUST_NOT_BE_COPIED" const callerInstructions = "CALLER_INSTRUCTIONS_MUST_NOT_BE_COPIED" const content = "MODEL_OUTPUT_SENTINEL" const reasoning = "MODEL_THINK_SENTINEL" body := []byte(`{"model":"served-responses","stream":false,"instructions":"` + callerInstructions + `","input":"` + callerInput + `"}`) _, rebuilder, source, ref := newOpenAIResumeRebuilderFixture(t, openAIRebuildEndpointResponses, body, 4096) source.recordText(content) source.recordReasoning(reasoning) directive, err := streamgate.NewRecoveryDirectiveContinuation(len(content), source.snapshotRef()) if err != nil { t.Fatalf("NewRecoveryDirectiveContinuation: %v", err) } draft, err := rebuilder.RebuildRequest(context.Background(), ref, mustOpenAIRecoveryPlan(t, "plan.responses.resume", streamgate.RecoveryStrategyContinuationRepair, directive)) if err != nil { t.Fatalf("RebuildRequest: %v", err) } lease, err := rebuilder.RebuiltStore().take(draft.RequestRef()) if err != nil { t.Fatalf("take: %v", err) } defer lease.release() got, err := lease.body() if err != nil { t.Fatalf("body: %v", err) } if bytes.Contains(got, []byte(callerInput)) || bytes.Contains(got, []byte(callerInstructions)) { t.Fatalf("responses resume copied caller input or instructions: %s", got) } var rebuilt struct { Model string `json:"model"` Stream bool `json:"stream"` Instructions string `json:"instructions"` Input []struct { Type string `json:"type"` Role string `json:"role"` Content []struct { Type string `json:"type"` Text string `json:"text"` } `json:"content"` } `json:"input"` } if err := json.Unmarshal(got, &rebuilt); err != nil { t.Fatalf("decode resume body: %v", err) } if rebuilt.Model != "served-responses" || rebuilt.Stream || rebuilt.Instructions != openAIRepeatResumeDirective || len(rebuilt.Input) != 2 { t.Fatalf("responses resume envelope = %#v", rebuilt) } if got := rebuilt.Input[0]; got.Type != "reasoning" || len(got.Content) != 1 || got.Content[0].Type != "reasoning_text" || got.Content[0].Text != reasoning { t.Fatalf("responses reasoning provenance = %#v", got) } if got := rebuilt.Input[1]; got.Type != "message" || got.Role != "assistant" || len(got.Content) != 1 || got.Content[0].Type != "output_text" || got.Content[0].Text != content { t.Fatalf("responses assistant provenance = %#v", got) } } func TestOpenAIResponsesRepeatResumeAdmissionAcceptsBuilderBody(t *testing.T) { const callerInput = "CALLER_RESPONSES_INPUT_MUST_NOT_BE_COPIED" const callerInstructions = "CALLER_INSTRUCTIONS_MUST_NOT_BE_COPIED" const content = "MODEL_OUTPUT_SENTINEL" const reasoning = "MODEL_THINK_SENTINEL" body := []byte(`{"model":"served-responses","stream":false,"instructions":"` + callerInstructions + `","input":"` + callerInput + `"}`) _, rebuilder, source, ref := newOpenAIResumeRebuilderFixture(t, openAIRebuildEndpointResponses, body, 4096) source.recordText(content) source.recordReasoning(reasoning) directive, err := streamgate.NewRecoveryDirectiveContinuation(len(content), source.snapshotRef()) if err != nil { t.Fatalf("NewRecoveryDirectiveContinuation: %v", err) } draft, err := rebuilder.RebuildRequest(context.Background(), ref, mustOpenAIRecoveryPlan(t, "plan.responses.resume.admission", streamgate.RecoveryStrategyContinuationRepair, directive)) if err != nil { t.Fatalf("RebuildRequest: %v", err) } lease, err := rebuilder.RebuiltStore().take(draft.RequestRef()) if err != nil { t.Fatalf("take: %v", err) } defer lease.release() resumeBody, err := lease.body() if err != nil { t.Fatalf("body: %v", err) } srv := NewServer(config.EdgeOpenAIConf{Adapter: "ollama", Target: "served", SessionID: "cli"}, &fakeRunService{}, nil) base := newTestRequestContext(t, routeDispatch{Adapter: "ollama", Target: "served", SessionID: "cli"}, body) base.endpoint = usageEndpointResponses requestCtx := &responsesRequestContext{openAIRequestContext: base, envelope: responsesEnvelope{Model: "served-responses"}} initial, err := srv.newResponsesDispatchContext(requestCtx, responsesRequest{Model: "served-responses", Input: json.RawMessage(`"initial"`)}) if err != nil { t.Fatalf("newResponsesDispatchContext: %v", err) } state := &openAIResponsesAttemptContext{} admission, err := newOpenAIResponsesRecoveryAdmissionBuilder(srv, initial, state)(context.Background(), streamgate.RebuiltRequest{}, resumeBody) if err != nil { t.Fatalf("resume admission: %v", err) } if admission.kind != openAIAdmissionRun || admission.run.Prompt == "" { t.Fatalf("resume admission = %#v, want normalized run", admission) } got := state.get() if got == nil { t.Fatal("resume admission did not retain dispatch context") } if strings.Contains(got.prompt, callerInput) || strings.Contains(got.prompt, callerInstructions) { t.Fatalf("resume prompt copied caller history: %q", got.prompt) } for _, want := range []string{openAIRepeatResumeDirective, content, reasoning} { if !strings.Contains(got.prompt, want) { t.Fatalf("resume prompt %q does not preserve %q", got.prompt, want) } } if gotValue, ok := got.input["responses_resume_content"].(string); !ok || gotValue != content { t.Fatalf("resume content input = %#v", got.input["responses_resume_content"]) } if gotValue, ok := got.input["responses_resume_reasoning"].(string); !ok || gotValue != reasoning { t.Fatalf("resume reasoning input = %#v", got.input["responses_resume_reasoning"]) } if err := decodeResponsesRequest(json.NewDecoder(bytes.NewReader(resumeBody)), &responsesRequest{}); err != nil { // Strict decoding is intentionally structural; public string-only input // is enforced at dispatch-context construction below. t.Fatalf("strict request decode unexpectedly failed: %v", err) } var public responsesRequest if err := decodeResponsesRequest(json.NewDecoder(bytes.NewReader(resumeBody)), &public); err != nil { t.Fatalf("public decode: %v", err) } if _, err := srv.newResponsesDispatchContext(requestCtx, public); err == nil { t.Fatal("public Responses array input was admitted") } } func TestOpenAIRecoverySourceStoreLifecycle(t *testing.T) { ingress, err := buildOpenAIIngressSnapshot(4096, []byte(`{"model":"m","messages":[]}`), json.RawMessage(`{"model":"m","messages":[]}`)) if err != nil { t.Fatalf("buildOpenAIIngressSnapshot: %v", err) } defer ingress.Close() source := newOpenAIRecoverySourceStore(ingress) source.recordText("first") source.recordReasoning("think") content, reasoning, err := source.consume(5) if err != nil || content != "first" || reasoning != "think" { t.Fatalf("first consume = (%q, %q, %v)", content, reasoning, err) } if _, _, err := source.consume(0); !errors.Is(err, errOpenAIRecoverySourceUnavailable) { t.Fatalf("second consume error = %v, want unavailable", err) } source.resetAttempt() source.recordText("next") content, reasoning, err = source.consume(4) if err != nil || content != "next" || reasoning != "" { t.Fatalf("reset consume = (%q, %q, %v)", content, reasoning, err) } source.close() if _, _, err := source.consume(0); !errors.Is(err, streamgate.ErrIngressSnapshotClosed) { t.Fatalf("closed consume error = %v, want ErrIngressSnapshotClosed", err) } } func TestRepeatGuardReasoningCursor(t *testing.T) { ingress, err := buildOpenAIIngressSnapshot( 4096, []byte(`{"model":"m","messages":[]}`), json.RawMessage(`{"model":"m","messages":[]}`), ) if err != nil { t.Fatalf("buildOpenAIIngressSnapshot: %v", err) } defer ingress.Close() source := newOpenAIRecoverySourceStore(ingress) source.recordText("content-prefix") source.recordReasoning("reasoning-prefix-repeated-tail") cursor := len("content-prefix") + 1 + len("reasoning-prefix-") content, reasoning, err := source.consume(cursor) if err != nil { t.Fatalf("consume reasoning cursor: %v", err) } if content != "content-prefix" || reasoning != "reasoning-prefix-" { t.Fatalf("reasoning cursor result = (%q, %q)", content, reasoning) } } func TestRepeatGuardKnownPrefixSuppression(t *testing.T) { ingress, err := buildOpenAIIngressSnapshot( 4096, []byte(`{"model":"m","messages":[]}`), json.RawMessage(`{"model":"m","messages":[]}`), ) if err != nil { t.Fatalf("buildOpenAIIngressSnapshot: %v", err) } defer ingress.Close() source := newOpenAIRecoverySourceStore(ingress) source.recordText("safe-prefix") if _, _, err := source.consume(len("safe-prefix")); err != nil { t.Fatalf("consume: %v", err) } source.resetAttempt() if got := source.suppressKnownPrefix(true, "safe-"); got != "" { t.Fatalf("first echoed fragment = %q, want suppressed", got) } if got := source.suppressKnownPrefix(true, "prefixnovel"); got != "novel" { t.Fatalf("second echoed fragment = %q, want novel suffix", got) } } func TestOpenAIRequestRebuilderRepeatResumeContextOverflow(t *testing.T) { body := []byte(`{"model":"served-chat","messages":[{"role":"user","content":"caller"}]}`) ingress, rebuilder, source, ref := newOpenAIResumeRebuilderFixture(t, openAIRebuildEndpointChat, body, 1) source.recordText("model output") directive, err := streamgate.NewRecoveryDirectiveContinuation(len("model output"), source.snapshotRef()) if err != nil { t.Fatalf("NewRecoveryDirectiveContinuation: %v", err) } _, err = rebuilder.RebuildRequest(context.Background(), ref, mustOpenAIRecoveryPlan(t, "plan.resume.overflow", streamgate.RecoveryStrategyContinuationRepair, directive)) if !errors.Is(err, errOpenAIRebuildContextOverflow) { t.Fatalf("RebuildRequest error = %v, want context overflow", err) } if got := len(rebuilder.RebuiltStore().leases); got != 0 { t.Fatalf("overflow retained %d dispatchable requests", got) } accessor, err := ingress.accessor() if err != nil { t.Fatalf("ingress accessor: %v", err) } if got := accessor.ReservedTempBytes(); got != 0 { t.Fatalf("overflow reserved bytes = %d, want 0", got) } } // TestOpenAIRepeatResumeContextOverflowPreservesBudget pins the host half of // the coordinator contract: continuation rebuild refusal creates no dispatchable // request lease. Recovery usage is consumed only immediately before an outbound // dispatcher call, so this fail-closed path leaves the fault budget at zero. func TestOpenAIRepeatResumeContextOverflowPreservesBudget(t *testing.T) { body := []byte(`{"model":"served-chat","messages":[{"role":"user","content":"caller"}]}`) _, rebuilder, source, ref := newOpenAIResumeRebuilderFixture(t, openAIRebuildEndpointChat, body, 1) source.recordText("model output") directive, err := streamgate.NewRecoveryDirectiveContinuation(len("model output"), source.snapshotRef()) if err != nil { t.Fatalf("NewRecoveryDirectiveContinuation: %v", err) } intent, err := streamgate.NewRecoveryIntent(streamgate.RecoveryStrategyContinuationRepair, directive, "repeat_detected", 10) if err != nil { t.Fatalf("NewRecoveryIntent: %v", err) } arbitration, err := streamgate.NewArbitrationResult(streamgate.ArbitrationActionRecover, streamgate.BaseDispositionTerminalErrorCandidate, "repeat_guard", "repeat_detected", &intent, nil) if err != nil { t.Fatalf("NewArbitrationResult: %v", err) } policy, err := streamgate.NewRecoveryPolicySnapshot(3, map[streamgate.RecoveryStrategy]int{streamgate.RecoveryStrategyContinuationRepair: 3}) if err != nil { t.Fatalf("NewRecoveryPolicySnapshot: %v", err) } usage, err := streamgate.NewRecoveryUsageSnapshot(0, nil) if err != nil { t.Fatalf("NewRecoveryUsageSnapshot: %v", err) } trace := &openAIResumeTrace{} controller := &openAIResumeTestController{trace: trace} binding, err := streamgate.NewAttemptBinding("attempt.initial", "served-chat", "provider", "normalized", openAIResumeTestEventSource{}, controller) if err != nil { t.Fatalf("NewAttemptBinding: %v", err) } dispatcher := &openAIResumeNoDispatch{} coordinator, err := streamgate.NewRecoveryCoordinator(streamgate.RecoveryCoordinatorOptions{ Policy: policy, Usage: usage, RequestSnapshot: ref, CurrentBinding: binding, Rebuilder: openAIResumeTracingRebuilder{inner: rebuilder, trace: trace}, Dispatcher: dispatcher, }) if err != nil { t.Fatalf("NewRecoveryCoordinator: %v", err) } result, err := coordinator.Execute(context.Background(), streamgate.RecoveryCycleInput{ Arbitration: arbitration, PlanID: "plan.resume.budget", IdempotencyKey: "plan.resume.budget:key", ConsumerID: "openai.edge", CommitState: streamgate.CommitStateStreamOpen, }) if !errors.Is(err, streamgate.ErrRecoveryRebuildFailed) { t.Fatalf("Execute error = %v, want ErrRecoveryRebuildFailed", err) } if controller.calls != 1 || dispatcher.calls != 0 { t.Fatalf("context refusal lifecycle = aborts:%d dispatches:%d, want 1/0", controller.calls, dispatcher.calls) } if got, want := trace.snapshot(), []string{"abort", "rebuild"}; !slices.Equal(got, want) { t.Fatalf("context refusal order = %v, want %v", got, want) } if got := coordinator.UsageSnapshot().RequestFaultRecoveries(); got != 0 { t.Fatalf("coordinator fault usage = %d, want 0", got) } if got := result.UsageSnapshot().RequestFaultRecoveries(); got != 0 { t.Fatalf("result fault usage = %d, want 0", got) } } func TestOpenAIRepeatResumeDoesNotUsePreparer(t *testing.T) { body := []byte(`{"model":"served-responses","input":"CALLER_INPUT_MUST_NOT_BE_COPIED"}`) _, rebuilder, source, ref := newOpenAIResumeRebuilderFixture(t, openAIRebuildEndpointResponses, body, 4096) source.recordText("safe prefix repeated tail") fake := &fakeRunService{eventRuns: []chan *iop.RunEvent{bufferedRunEvents(&iop.RunEvent{Type: "complete"})}} srv := NewServer(config.EdgeOpenAIConf{Adapter: "ollama", Target: "served", SessionID: "cli"}, fake, nil) base := newTestRequestContext(t, routeDispatch{Adapter: "ollama", Target: "served", SessionID: "cli"}, body) base.endpoint = usageEndpointResponses requestCtx := &responsesRequestContext{openAIRequestContext: base, envelope: responsesEnvelope{Model: "served-responses"}} dc, err := srv.newResponsesDispatchContext(requestCtx, responsesRequest{Model: "served-responses", Input: json.RawMessage(`"CALLER_INPUT_MUST_NOT_BE_COPIED"`)}) if err != nil { t.Fatalf("newResponsesDispatchContext: %v", err) } state := &openAIResponsesAttemptContext{} dispatcher, err := newOpenAIAttemptDispatcher(srv.service, rebuilder.RebuiltStore(), newOpenAIResponsesRecoveryAdmissionBuilder(srv, dc, state), func(transport openAIAttemptTransport) (streamgate.NormalizedEventSource, error) { if transport.run == nil || state.get() == nil { return nil, errors.New("resume was not admitted") } return newOpenAIResponsesEventSource(state.get(), transport.run, &openAIResponsesResultHolder{}, nil), nil }) if err != nil { t.Fatalf("newOpenAIAttemptDispatcher: %v", err) } directive, err := streamgate.NewRecoveryDirectiveContinuation(len("safe prefix "), source.snapshotRef()) if err != nil { t.Fatalf("NewRecoveryDirectiveContinuation: %v", err) } intent, err := streamgate.NewRecoveryIntent(streamgate.RecoveryStrategyContinuationRepair, directive, "repeat_detected", 10) if err != nil { t.Fatalf("NewRecoveryIntent: %v", err) } arbitration, err := streamgate.NewArbitrationResult(streamgate.ArbitrationActionRecover, streamgate.BaseDispositionTerminalErrorCandidate, "repeat_guard", "repeat_detected", &intent, nil) if err != nil { t.Fatalf("NewArbitrationResult: %v", err) } policy, err := streamgate.NewRecoveryPolicySnapshot(3, map[streamgate.RecoveryStrategy]int{streamgate.RecoveryStrategyContinuationRepair: 3}) if err != nil { t.Fatalf("NewRecoveryPolicySnapshot: %v", err) } usage, err := streamgate.NewRecoveryUsageSnapshot(0, nil) if err != nil { t.Fatalf("NewRecoveryUsageSnapshot: %v", err) } trace := &openAIResumeTrace{} controller := &openAIResumeTestController{trace: trace} binding, err := streamgate.NewAttemptBinding("attempt.initial", "served-responses", "provider", "normalized", openAIResumeTestEventSource{}, controller) if err != nil { t.Fatalf("NewAttemptBinding: %v", err) } preparer := &openAIResumeRecordingPreparer{} coordinator, err := streamgate.NewRecoveryCoordinator(streamgate.RecoveryCoordinatorOptions{ Policy: policy, Usage: usage, RequestSnapshot: ref, CurrentBinding: binding, Rebuilder: openAIResumeTracingRebuilder{inner: rebuilder, trace: trace}, Dispatcher: openAIResumeTracingDispatcher{inner: dispatcher, trace: trace}, Preparers: map[string]streamgate.RecoveryPlanPreparer{"recording": preparer}, PreparationTimeout: time.Second, }) if err != nil { t.Fatalf("NewRecoveryCoordinator: %v", err) } result, err := coordinator.Execute(context.Background(), streamgate.RecoveryCycleInput{ Arbitration: arbitration, PlanID: "plan.resume.no-preparer", IdempotencyKey: "plan.resume.no-preparer:key", ConsumerID: "openai.edge", CommitState: streamgate.CommitStateStreamOpen, }) if err != nil { t.Fatalf("Execute: %v", err) } if preparer.calls != 0 { t.Fatalf("continuation invoked preparer %d times, want 0", preparer.calls) } if controller.calls != 1 || len(fake.reqsSnapshot()) != 1 { t.Fatalf("lifecycle aborts=%d dispatches=%d, want one abort then one dispatch", controller.calls, len(fake.reqsSnapshot())) } if got, want := trace.snapshot(), []string{"abort", "rebuild", "dispatch"}; !slices.Equal(got, want) { t.Fatalf("recovery order = %v, want %v", got, want) } if got := result.UsageSnapshot().RequestFaultRecoveries(); got != 1 { t.Fatalf("fault usage after actual dispatch = %d, want 1", got) } } func TestOpenAIRequestRebuilderExactByteIdentity(t *testing.T) { body := []byte("{\n \"unknown\" : [1, 2], \"model\" : \"alias\", \"messages\" : []\n}\n") _, rebuilder, ref := newOpenAIRebuilderFixture(t, openAIRebuildEndpointChat, body, 4096) directive, err := streamgate.NewRecoveryDirectiveExact(ref.SnapshotRef()) if err != nil { t.Fatalf("NewRecoveryDirectiveExact: %v", err) } plan := mustOpenAIRecoveryPlan(t, "plan.exact", streamgate.RecoveryStrategyExactReplay, directive) draft, err := rebuilder.RebuildRequest(context.Background(), ref, plan) if err != nil { t.Fatalf("RebuildRequest: %v", err) } lease, err := rebuilder.RebuiltStore().take(draft.RequestRef()) if err != nil { t.Fatalf("take: %v", err) } defer lease.release() got, err := lease.body() if err != nil { t.Fatalf("body: %v", err) } if !bytes.Equal(got, body) { t.Fatalf("exact rebuild changed bytes:\n%s", got) } } func TestOpenAIRequestRebuilderContinuationPreservesNonTargetBytes(t *testing.T) { body := []byte("{\n \"unknown\" : { \"x\": 1 }, \"model\" : \"alias\", \"messages\" : [ {\"role\":\"user\",\"content\":\"old\"} ], \"tail\" : true\n}") _, rebuilder, ref := newOpenAIRebuilderFixture(t, openAIRebuildEndpointChat, body, 8192) patch := json.RawMessage(`[{"role":"assistant","content":"safe prefix"},{"role":"user","content":"continue"}]`) if err := rebuilder.PatchStore().PutContinuation("snapshot.continue", 17, patch); err != nil { t.Fatalf("PutContinuation: %v", err) } directive, _ := streamgate.NewRecoveryDirectiveContinuation(17, "snapshot.continue") plan := mustOpenAIRecoveryPlan(t, "plan.continue", streamgate.RecoveryStrategyContinuationRepair, directive) draft, err := rebuilder.RebuildRequest(context.Background(), ref, plan) if err != nil { t.Fatalf("RebuildRequest: %v", err) } lease, _ := rebuilder.RebuiltStore().take(draft.RequestRef()) defer lease.release() got, _ := lease.body() want := bytes.Replace(body, []byte(`[ {"role":"user","content":"old"} ]`), patch, 1) if !bytes.Equal(got, want) { t.Fatalf("continuation rebuild changed non-target bytes:\n got=%s\nwant=%s", got, want) } if draft.RetainedBytes() > draft.PeakBytes() || draft.PeakBytes() > draft.MaxBytes() { t.Fatalf("invalid draft memory accounting: retained=%d peak=%d max=%d", draft.RetainedBytes(), draft.PeakBytes(), draft.MaxBytes()) } } func TestOpenAIRequestRebuilderResponsesSchemaPatch(t *testing.T) { body := []byte(`{ "model" : "alias", "input" : { "old" : true }, "custom" : [3, 2, 1] }`) _, rebuilder, ref := newOpenAIRebuilderFixture(t, openAIRebuildEndpointResponses, body, 4096) patch := json.RawMessage(`[{"role":"user","content":"fixed"}]`) if err := rebuilder.PatchStore().PutSchema("schema.responses", "patch.input", patch); err != nil { t.Fatalf("PutSchema: %v", err) } directive, _ := streamgate.NewRecoveryDirectiveSchema("schema.responses", "patch.input") plan := mustOpenAIRecoveryPlan(t, "plan.schema", streamgate.RecoveryStrategySchemaRepair, directive) draft, err := rebuilder.RebuildRequest(context.Background(), ref, plan) if err != nil { t.Fatalf("RebuildRequest: %v", err) } lease, _ := rebuilder.RebuiltStore().take(draft.RequestRef()) defer lease.release() got, _ := lease.body() want := bytes.Replace(body, []byte(`{ "old" : true }`), patch, 1) if !bytes.Equal(got, want) { t.Fatalf("schema rebuild changed non-target bytes:\n got=%s\nwant=%s", got, want) } } // TestOpenAIResponsesCodecEndpointDistinctUnknownPreservation proves the S18 // Responses lossless codec: a Responses rebuild patches the top-level "input" // field while preserving unknown top-level fields and unknown/encrypted input // items, and the same patch shape on the Chat endpoint targets "messages" // instead (endpoint-specific, no shared parser). func TestOpenAIResponsesCodecEndpointDistinctUnknownPreservation(t *testing.T) { respBody := []byte(`{ "model" : "alias", "input" : [ {"type":"reasoning","encrypted":"zzz"}, {"type":"message","role":"user","content":"old"} ], "custom" : [3, 2, 1], "future" : {"x":true} }`) _, respRebuilder, respRef := newOpenAIRebuilderFixture(t, openAIRebuildEndpointResponses, respBody, 8192) respPatch := json.RawMessage(`[{"type":"message","role":"user","content":"fixed"}]`) if err := respRebuilder.PatchStore().PutContinuation("snapshot.resp", 5, respPatch); err != nil { t.Fatalf("PutContinuation(responses): %v", err) } respDirective, _ := streamgate.NewRecoveryDirectiveContinuation(5, "snapshot.resp") respPlan := mustOpenAIRecoveryPlan(t, "plan.resp", streamgate.RecoveryStrategyContinuationRepair, respDirective) respDraft, err := respRebuilder.RebuildRequest(context.Background(), respRef, respPlan) if err != nil { t.Fatalf("RebuildRequest(responses): %v", err) } respLease, _ := respRebuilder.RebuiltStore().take(respDraft.RequestRef()) defer respLease.release() respGot, _ := respLease.body() wantResp := bytes.Replace(respBody, []byte(`[ {"type":"reasoning","encrypted":"zzz"}, {"type":"message","role":"user","content":"old"} ]`), respPatch, 1) if !bytes.Equal(respGot, wantResp) { t.Fatalf("responses rebuild changed non-target bytes:\n got=%s\nwant=%s", respGot, wantResp) } for _, keep := range []string{`"custom" : [3, 2, 1]`, `"future" : {"x":true}`} { if !bytes.Contains(respGot, []byte(keep)) { t.Errorf("responses rebuild dropped unknown top-level field %q", keep) } } // Same patch shape on the Chat endpoint targets "messages", not "input"; // a Chat body's "input" field is an unknown field and must be preserved. chatBody := []byte(`{ "model" : "alias", "messages" : [ {"role":"user","content":"old"} ], "input" : "unknown-to-chat" }`) _, chatRebuilder, chatRef := newOpenAIRebuilderFixture(t, openAIRebuildEndpointChat, chatBody, 8192) chatPatch := json.RawMessage(`[{"role":"user","content":"fixed"}]`) if err := chatRebuilder.PatchStore().PutContinuation("snapshot.chat", 5, chatPatch); err != nil { t.Fatalf("PutContinuation(chat): %v", err) } chatDirective, _ := streamgate.NewRecoveryDirectiveContinuation(5, "snapshot.chat") chatPlan := mustOpenAIRecoveryPlan(t, "plan.chat", streamgate.RecoveryStrategyContinuationRepair, chatDirective) chatDraft, err := chatRebuilder.RebuildRequest(context.Background(), chatRef, chatPlan) if err != nil { t.Fatalf("RebuildRequest(chat): %v", err) } chatLease, _ := chatRebuilder.RebuiltStore().take(chatDraft.RequestRef()) defer chatLease.release() chatGot, _ := chatLease.body() wantChat := bytes.Replace(chatBody, []byte(`[ {"role":"user","content":"old"} ]`), chatPatch, 1) if !bytes.Equal(chatGot, wantChat) { t.Fatalf("chat rebuild changed non-target bytes:\n got=%s\nwant=%s", chatGot, wantChat) } if !bytes.Contains(chatGot, []byte(`"input" : "unknown-to-chat"`)) { t.Errorf("chat rebuild dropped its unknown top-level \"input\" field") } } func TestOpenAIRequestRebuilderPatchPlusOutputPeakOverflow(t *testing.T) { body := []byte(`{"model":"m","messages":[]}`) patch := json.RawMessage(`[{"role":"user","content":"larger"}]`) patchPlan, err := planTopLevelJSONPatches(body, []topLevelJSONPatch{{name: "messages", value: patch}}) if err != nil { t.Fatalf("planTopLevelJSONPatches: %v", err) } maxBytes := int64(len(body) + len(patch) + patchPlan.outputSize - 1) ingress, rebuilder, ref := newOpenAIRebuilderFixture(t, openAIRebuildEndpointChat, body, maxBytes) if err := rebuilder.PatchStore().PutSchema("schema.chat", "patch.messages", patch); err != nil { t.Fatalf("PutSchema: %v", err) } directive, _ := streamgate.NewRecoveryDirectiveSchema("schema.chat", "patch.messages") plan := mustOpenAIRecoveryPlan(t, "plan.overflow", streamgate.RecoveryStrategySchemaRepair, directive) _, err = rebuilder.RebuildRequest(context.Background(), ref, plan) if !errors.Is(err, streamgate.ErrIngressSnapshotRebuildOverflow) { t.Fatalf("error = %v, want rebuild overflow", err) } if len(rebuilder.RebuiltStore().leases) != 0 { t.Fatal("overflow retained a dispatchable request lease") } if len(rebuilder.PatchStore().schema) != 0 { t.Fatal("overflow retained a one-shot patch") } accessor, accessErr := ingress.accessor() if accessErr != nil { t.Fatalf("accessor after pre-allocation overflow: %v", accessErr) } if got := accessor.ReservedTempBytes(); got != 0 { t.Fatalf("reserved bytes after overflow = %d, want 0", got) } } func TestProviderRequestPatchesPreserveUnknownOrderAndWhitespace(t *testing.T) { chat := []byte(`{ "z" : 1, "model" : "alias", "max_completion_tokens" : 7, "unknown" : { "a" : 2 } }`) maxTokens := 9 gotChat, err := rewriteChatCompletionModel(chat, "served", chatCompletionRequest{MaxTokens: &maxTokens}) if err != nil { t.Fatalf("rewriteChatCompletionModel: %v", err) } if !bytes.Contains(gotChat, []byte(`"z" : 1`)) || !bytes.Contains(gotChat, []byte(`"unknown" : { "a" : 2 }`)) { t.Fatalf("chat unknown byte ranges changed: %s", gotChat) } if bytes.Contains(gotChat, []byte("max_completion_tokens")) || !bytes.Contains(gotChat, []byte(`"max_tokens":9`)) { t.Fatalf("chat policy patch mismatch: %s", gotChat) } responses := []byte("{\n \"input\" : \"hello\", \"model\" : \"alias\", \"future\" : {\"x\":true}\n}") gotResponses, err := rewriteResponsesModel(responses, "served") if err != nil { t.Fatalf("rewriteResponsesModel: %v", err) } wantResponses := bytes.Replace(responses, []byte(`"alias"`), []byte(`"served"`), 1) if !bytes.Equal(gotResponses, wantResponses) { t.Fatalf("responses non-target bytes changed:\n%s", gotResponses) } } func TestOpenAIRequestRebuilderRejectsReferenceMismatch(t *testing.T) { body := []byte(`{"model":"m","messages":[]}`) _, rebuilder, ref := newOpenAIRebuilderFixture(t, openAIRebuildEndpointChat, body, 1024) directive, _ := streamgate.NewRecoveryDirectiveExact(ref.SnapshotRef()) plan := mustOpenAIRecoveryPlan(t, "plan.mismatch", streamgate.RecoveryStrategyExactReplay, directive) wrong, err := streamgate.NewRecoveryRequestSnapshotRef("openai.ingress.wrong", ref.RetainedBytes(), ref.PeakBytes(), ref.MaxBytes()) if err != nil { t.Fatalf("NewRecoveryRequestSnapshotRef: %v", err) } if _, err := rebuilder.RebuildRequest(context.Background(), wrong, plan); err == nil { t.Fatal("mismatched snapshot reference was accepted") } } func Example_openAIRequestRebuilder() { fmt.Println(openAIRebuildFamily) // Output: openai.json } type cancelAfterFirstErrContext struct { context.Context calls int } func (c *cancelAfterFirstErrContext) Err() error { c.calls++ if c.calls > 1 { return context.Canceled } return nil } func TestOpenAIRequestRebuilderActualOwnedPeakBoundaries(t *testing.T) { body := []byte(`{"model":"alias","input":{"old":true},"custom":7}`) patch := json.RawMessage(`[{"role":"user","content":"fixed"}]`) patchPlan, err := planTopLevelJSONPatches(body, []topLevelJSONPatch{{name: "input", value: patch}}) if err != nil { t.Fatalf("planTopLevelJSONPatches: %v", err) } maxBytes := int64(len(body) + len(patch) + patchPlan.outputSize) ingress, rebuilder, ref := newOpenAIRebuilderFixture(t, openAIRebuildEndpointResponses, body, maxBytes) if err := rebuilder.PatchStore().PutSchema("schema.actual", "patch.input", patch); err != nil { t.Fatalf("PutSchema: %v", err) } directive, _ := streamgate.NewRecoveryDirectiveSchema("schema.actual", "patch.input") plan := mustOpenAIRecoveryPlan(t, "plan.actual", streamgate.RecoveryStrategySchemaRepair, directive) draft, err := rebuilder.RebuildRequest(context.Background(), ref, plan) if err != nil { t.Fatalf("RebuildRequest at exact owned peak: %v", err) } wantRetained := uint64(len(body) + patchPlan.outputSize) if draft.RetainedBytes() != wantRetained || draft.PeakBytes() != uint64(maxBytes) || draft.MaxBytes() != uint64(maxBytes) { t.Fatalf("draft accounting = retained:%d peak:%d max:%d, want %d/%d/%d", draft.RetainedBytes(), draft.PeakBytes(), draft.MaxBytes(), wantRetained, maxBytes, maxBytes) } accessor, err := ingress.accessor() if err != nil { t.Fatalf("accessor: %v", err) } if got := accessor.ReservedTempBytes(); got != 0 { t.Fatalf("patch reservation after rebuild = %d, want 0", got) } lease, err := rebuilder.RebuiltStore().take(draft.RequestRef()) if err != nil { t.Fatalf("take rebuilt lease: %v", err) } got, err := lease.body() if err != nil { t.Fatalf("lease body: %v", err) } typedAlias, err := lease.rebuilt.Accessor().TypedViewAlias(openAIRebuiltBodyViewName) if err != nil { t.Fatalf("TypedViewAlias: %v", err) } if len(got) == 0 || &got[0] != &typedAlias[0] { t.Fatal("dispatch lease did not retain the committed owned output alias") } guard := lease.guard lease.release() if !lease.isReleased() || !guard.IsReleased() { t.Fatal("rebuilt output lease did not release snapshot and guard") } } func TestOpenAIRequestRebuilderPatchStoreBoundedOneShotRelease(t *testing.T) { body := []byte(`{"model":"m","messages":[]}`) patch := json.RawMessage(`["owned"]`) ingress, rebuilder, ref := newOpenAIRebuilderFixture(t, openAIRebuildEndpointChat, body, int64(len(body)+128)) store := rebuilder.PatchStore() if err := store.PutContinuation("snapshot.once", 9, patch); err != nil { t.Fatalf("PutContinuation: %v", err) } accessor, _ := ingress.accessor() if got := accessor.ReservedTempBytes(); got != int64(len(patch)) { t.Fatalf("reserved patch bytes = %d, want %d", got, len(patch)) } store.mu.Lock() stored := store.continuation["snapshot.once"] store.mu.Unlock() if stored == nil || len(stored.value) == 0 || &stored.value[0] != &patch[0] { t.Fatal("json.RawMessage patch was copied instead of ownership-transferred") } if err := store.PutContinuation("snapshot.once", 9, json.RawMessage(`["duplicate"]`)); !errors.Is(err, errOpenAIRecoveryPatchDuplicate) { t.Fatalf("duplicate PutContinuation = %v, want duplicate", err) } if got := accessor.ReservedTempBytes(); got != int64(len(patch)) { t.Fatalf("duplicate changed reservation to %d", got) } entry, err := store.takeContinuation("snapshot.once", 9) if err != nil { t.Fatalf("takeContinuation: %v", err) } if _, err := store.takeContinuation("snapshot.once", 9); err == nil { t.Fatal("one-shot continuation patch was available twice") } entry.release() if got := accessor.ReservedTempBytes(); got != 0 { t.Fatalf("reservation after one-shot release = %d, want 0", got) } cancelPatch := json.RawMessage(`{"cancelled":true}`) if err := store.PutSchema("schema.cancel", "patch.cancel", cancelPatch); err != nil { t.Fatalf("PutSchema(cancel): %v", err) } directive, _ := streamgate.NewRecoveryDirectiveSchema("schema.cancel", "patch.cancel") plan := mustOpenAIRecoveryPlan(t, "plan.cancel", streamgate.RecoveryStrategySchemaRepair, directive) cancelCtx := &cancelAfterFirstErrContext{Context: context.Background()} if _, err := rebuilder.RebuildRequest(cancelCtx, ref, plan); !errors.Is(err, context.Canceled) { t.Fatalf("cancelled RebuildRequest = %v, want context canceled", err) } if got := accessor.ReservedTempBytes(); got != 0 { t.Fatalf("reservation after cancellation = %d, want 0", got) } if len(store.schema) != 0 || len(rebuilder.RebuiltStore().leases) != 0 { t.Fatal("cancellation retained a patch or dispatch lease") } closePatch := json.RawMessage(`{"close":true}`) if err := store.PutSchema("schema.close", "patch.close", closePatch); err != nil { t.Fatalf("PutSchema(close): %v", err) } rebuilder.Close() if got := accessor.ReservedTempBytes(); got != 0 { t.Fatalf("reservation after Close = %d, want 0", got) } if !store.closed || len(store.schema) != 0 || len(store.continuation) != 0 { t.Fatal("Close did not empty and close the patch store") } oversizedIngress, oversizedRebuilder, _ := newOpenAIRebuilderFixture(t, openAIRebuildEndpointChat, body, int64(len(body)+2)) if err := oversizedRebuilder.PatchStore().PutSchema("schema.large", "patch.large", json.RawMessage(`{"large":true}`)); !errors.Is(err, streamgate.ErrIngressSnapshotRebuildOverflow) { t.Fatalf("oversized PutSchema = %v, want rebuild overflow", err) } oversizedAccessor, _ := oversizedIngress.accessor() if got := oversizedAccessor.ReservedTempBytes(); got != 0 { t.Fatalf("oversized patch left %d reserved bytes", got) } if len(oversizedRebuilder.PatchStore().schema) != 0 { t.Fatal("oversized patch entered the store") } } func TestOpenAIRequestRebuilderCloseIsTerminal(t *testing.T) { body := []byte(`{"model":"alias","messages":[]}`) ingress, rebuilder, ref := newOpenAIRebuilderFixture(t, openAIRebuildEndpointChat, body, 4096) directive, _ := streamgate.NewRecoveryDirectiveExact(ref.SnapshotRef()) plan := mustOpenAIRecoveryPlan(t, "plan.terminal", streamgate.RecoveryStrategyExactReplay, directive) if _, err := rebuilder.RebuildRequest(context.Background(), ref, plan); err != nil { t.Fatalf("pre-close RebuildRequest: %v", err) } rebuilder.Close() _, err := rebuilder.RebuildRequest(context.Background(), ref, plan) if !errors.Is(err, streamgate.ErrIngressSnapshotClosed) { t.Fatalf("post-close RebuildRequest = %v, want ErrIngressSnapshotClosed", err) } rebuilder.mu.Lock() closed := rebuilder.closed rebuilder.mu.Unlock() if !closed { t.Fatal("rebuilder closed flag not set") } rebuilt := rebuilder.RebuiltStore() rebuilt.mu.Lock() storeClosed := rebuilt.closed storeClosed2 := rebuilt.closed storeLeases := rebuilt.leases rebuilt.mu.Unlock() if !storeClosed { t.Fatal("rebuilt store closed flag not set") } if storeClosed2 != storeClosed { t.Fatal("rebuilt store closed flag inconsistent") } if len(storeLeases) != 0 { t.Fatalf("rebuilt store leases not empty after close: %d", len(storeLeases)) } accessor, _ := ingress.accessor() if got := accessor.ReservedTempBytes(); got != 0 { t.Fatalf("reserved bytes after Close = %d, want 0", got) } if _, err := rebuilder.RebuildRequest(context.Background(), ref, plan); !errors.Is(err, streamgate.ErrIngressSnapshotClosed) { t.Fatalf("second post-close RebuildRequest = %v, want ErrIngressSnapshotClosed", err) } } func TestOpenAIRequestRebuilderNilReceiver(t *testing.T) { var rebuilder *openAIRequestRebuilder rebuilder.Close() body := []byte(`{"model":"alias","messages":[]}`) _, _, ref := newOpenAIRebuilderFixture(t, openAIRebuildEndpointChat, body, 4096) directive, _ := streamgate.NewRecoveryDirectiveExact(ref.SnapshotRef()) plan := mustOpenAIRecoveryPlan(t, "plan.nil", streamgate.RecoveryStrategyExactReplay, directive) _, err := rebuilder.RebuildRequest(context.Background(), ref, plan) if !errors.Is(err, streamgate.ErrIngressSnapshotClosed) { t.Fatalf("nil RebuildRequest = %v, want ErrIngressSnapshotClosed", err) } } func TestOpenAIRequestRebuilderCloseWaitsForInFlightPatchedRebuild(t *testing.T) { body := []byte(`{"model":"alias","messages":[{"role":"user","content":"hello"}]}`) t.Run("continuation", func(t *testing.T) { ingress, rebuilder, ref := newOpenAIRebuilderFixture(t, openAIRebuildEndpointChat, body, 8192) patch := json.RawMessage(`[{"role":"assistant","content":"response"},{"role":"user","content":"continue"}]`) if err := rebuilder.PatchStore().PutContinuation("snap.cont", 1, patch); err != nil { t.Fatalf("PutContinuation: %v", err) } directive, _ := streamgate.NewRecoveryDirectiveContinuation(1, "snap.cont") plan := mustOpenAIRecoveryPlan(t, "plan.cont", streamgate.RecoveryStrategyContinuationRepair, directive) // Use a context that cancels after the first Err() call so the rebuild // takes the patch then returns context.Canceled, exercising the full // in-flight → deregister → Close unblock path deterministically. ctx := &cancelAfterFirstErrContext{Context: context.Background()} rebuildDone := make(chan struct{}) go func() { defer close(rebuildDone) _, _ = rebuilder.RebuildRequest(ctx, ref, plan) }() <-rebuildDone closeDone := make(chan struct{}) go func() { rebuilder.Close() close(closeDone) }() <-closeDone accessor, _ := ingress.accessor() if got := accessor.ReservedTempBytes(); got != 0 { t.Fatalf("reserved bytes after Close = %d, want 0", got) } rebuilt := rebuilder.RebuiltStore() rebuilt.mu.Lock() if !rebuilt.closed || len(rebuilt.leases) != 0 { t.Fatal("rebuilt store not closed or has leases after Close") } rebuilt.mu.Unlock() patches := rebuilder.PatchStore() patches.mu.Lock() if !patches.closed || len(patches.continuation) != 0 || len(patches.schema) != 0 { t.Fatal("patch store not closed or has entries after Close") } patches.mu.Unlock() }) t.Run("schema", func(t *testing.T) { ingress, rebuilder, ref := newOpenAIRebuilderFixture(t, openAIRebuildEndpointResponses, body, 8192) patch := json.RawMessage(`[{"role":"user","content":"fixed"}]`) if err := rebuilder.PatchStore().PutSchema("schema.resp", "patch.input", patch); err != nil { t.Fatalf("PutSchema: %v", err) } directive, _ := streamgate.NewRecoveryDirectiveSchema("schema.resp", "patch.input") plan := mustOpenAIRecoveryPlan(t, "plan.schema", streamgate.RecoveryStrategySchemaRepair, directive) ctx := &cancelAfterFirstErrContext{Context: context.Background()} rebuildDone := make(chan struct{}) go func() { defer close(rebuildDone) _, _ = rebuilder.RebuildRequest(ctx, ref, plan) }() <-rebuildDone closeDone := make(chan struct{}) go func() { rebuilder.Close() close(closeDone) }() <-closeDone accessor, _ := ingress.accessor() if got := accessor.ReservedTempBytes(); got != 0 { t.Fatalf("reserved bytes after Close = %d, want 0", got) } }) t.Run("duplicateClose", func(t *testing.T) { _, rebuilder, _ := newOpenAIRebuilderFixture(t, openAIRebuildEndpointChat, body, 8192) rebuilder.Close() rebuilder.Close() }) } func TestOpenAIRequestRebuilderConcurrentCloseWaitsForStoreDrain(t *testing.T) { body := []byte(`{"model":"alias","messages":[{"role":"user","content":"hello"}]}`) t.Run("continuation", func(t *testing.T) { ingress, rebuilder, ref := newOpenAIRebuilderFixture(t, openAIRebuildEndpointChat, body, 8192) patch := json.RawMessage(`[{"role":"assistant","content":"response"},{"role":"user","content":"continue"}]`) if err := rebuilder.PatchStore().PutContinuation("snap.cont.det", 1, patch); err != nil { t.Fatalf("PutContinuation: %v", err) } directive, _ := streamgate.NewRecoveryDirectiveContinuation(1, "snap.cont.det") plan := mustOpenAIRecoveryPlan(t, "plan.cont.det", streamgate.RecoveryStrategyContinuationRepair, directive) entered := make(chan struct{}) blockRebuild := make(chan struct{}) ctx := &blockingRebuildContext{Context: context.Background(), blocked: blockRebuild, entered: entered} rebuildDone := make(chan struct{}) go func() { defer close(rebuildDone) _, _ = rebuilder.RebuildRequest(ctx, ref, plan) }() // Wait for rebuild to take the patch and be blocked in ctx.Err(). <-entered // Launch two Close callers that must both register as waiters. close1Done := make(chan struct{}) close2Done := make(chan struct{}) go func() { defer close(close1Done) rebuilder.Close() }() go func() { defer close(close2Done) rebuilder.Close() }() // Wait for both Closers to actually register as close waiters under // the mutex. The helper observes cond.Broadcast, so no time.Sleep is // needed to estimate scheduler timing. waitForOpenAIRebuilderCloseWaiters(t, rebuilder, 2) // Verify neither Close has returned yet (done channels not closed). select { case <-close1Done: t.Fatal("close1 returned before the rebuild completed") default: } select { case <-close2Done: t.Fatal("close2 returned before the rebuild completed") default: } // Unblock the rebuild so all waiters can proceed. close(blockRebuild) <-rebuildDone <-close1Done <-close2Done // Verify stores are drained and reservation is 0. accessor, _ := ingress.accessor() if got := accessor.ReservedTempBytes(); got != 0 { t.Fatalf("reserved bytes after Close = %d, want 0", got) } rebuilt := rebuilder.RebuiltStore() rebuilt.mu.Lock() if !rebuilt.closed || len(rebuilt.leases) != 0 { t.Fatal("rebuilt store not closed or has leases after Close") } rebuilt.mu.Unlock() patches := rebuilder.PatchStore() patches.mu.Lock() if !patches.closed || len(patches.continuation) != 0 || len(patches.schema) != 0 { t.Fatal("patch store not closed or has entries after Close") } patches.mu.Unlock() }) t.Run("schema", func(t *testing.T) { ingress, rebuilder, ref := newOpenAIRebuilderFixture(t, openAIRebuildEndpointResponses, body, 8192) patch := json.RawMessage(`[{"role":"user","content":"fixed"}]`) if err := rebuilder.PatchStore().PutSchema("schema.resp.det", "patch.input", patch); err != nil { t.Fatalf("PutSchema: %v", err) } directive, _ := streamgate.NewRecoveryDirectiveSchema("schema.resp.det", "patch.input") plan := mustOpenAIRecoveryPlan(t, "plan.schema.det", streamgate.RecoveryStrategySchemaRepair, directive) entered := make(chan struct{}) blockRebuild := make(chan struct{}) ctx := &blockingRebuildContext{Context: context.Background(), blocked: blockRebuild, entered: entered} rebuildDone := make(chan struct{}) go func() { defer close(rebuildDone) _, _ = rebuilder.RebuildRequest(ctx, ref, plan) }() <-entered close1Done := make(chan struct{}) close2Done := make(chan struct{}) go func() { defer close(close1Done) rebuilder.Close() }() go func() { defer close(close2Done) rebuilder.Close() }() waitForOpenAIRebuilderCloseWaiters(t, rebuilder, 2) select { case <-close1Done: t.Fatal("close1 returned before the rebuild completed") default: } select { case <-close2Done: t.Fatal("close2 returned before the rebuild completed") default: } close(blockRebuild) <-rebuildDone <-close1Done <-close2Done accessor, _ := ingress.accessor() if got := accessor.ReservedTempBytes(); got != 0 { t.Fatalf("reserved bytes after Close = %d, want 0", got) } }) t.Run("nilReceiver", func(t *testing.T) { var rebuilder *openAIRequestRebuilder rebuilder.Close() rebuilder.Close() }) } // waitForOpenAIRebuilderCloseWaiters waits until the rebuilder has exactly // n goroutines registered as close waiters. It observes actual registration // under the mutex rather than estimating via time.Sleep. func waitForOpenAIRebuilderCloseWaiters(t *testing.T, rebuilder *openAIRequestRebuilder, n int) { t.Helper() rebuilder.mu.Lock() defer rebuilder.mu.Unlock() for rebuilder.closeWaiters != n { rebuilder.cond.Wait() } } type blockingRebuildContext struct { context.Context calls int blocked chan struct{} entered chan struct{} once sync.Once } func (c *blockingRebuildContext) Err() error { c.calls++ if c.calls == 1 { return nil // Allow rebuild to start, increment inFlight, and take the patch } // Patch has been taken; signal entry then block to keep rebuild in-flight. c.once.Do(func() { close(c.entered) }) <-c.blocked return c.Context.Err() } func TestOpenAIProviderBodyLeaseRelease(t *testing.T) { body := []byte(`{"model":"alias","input":"hello","custom":true}`) modelJSON, _ := json.Marshal("served") patchPlan, err := planTopLevelJSONPatches(body, []topLevelJSONPatch{{name: "model", value: modelJSON}}) if err != nil { t.Fatalf("planTopLevelJSONPatches: %v", err) } ingress, err := buildOpenAIIngressSnapshot(int64(len(body)+patchPlan.outputSize), body, json.RawMessage(body)) if err != nil { t.Fatalf("buildOpenAIIngressSnapshot: %v", err) } defer ingress.Close() builder := newOpenAIProviderBodyBuilder(func(target string) (*openAIRebuiltLease, error) { return rewriteResponsesModelFromIngress(ingress, target) }) got, err := builder.BuildBody("served") if err != nil { t.Fatalf("BuildBody: %v", err) } if !bytes.Contains(got, []byte(`"model":"served"`)) || !bytes.Contains(got, []byte(`"custom":true`)) { t.Fatalf("provider body rewrite mismatch: %s", got) } builder.mu.Lock() lease := builder.lease builder.mu.Unlock() if lease == nil || lease.guard == nil || lease.isReleased() || lease.guard.IsReleased() { t.Fatal("provider body lease was not live during synchronous submission") } guard := lease.guard builder.Close() if !lease.isReleased() || !guard.IsReleased() { t.Fatal("provider body lease was not released after synchronous submission") } if _, err := builder.BuildBody("served-again"); err == nil { t.Fatal("provider body builder allowed a second build") } } // TestOpenAIRequestRebuilderCloseWaitsForInFlightPatchedRebuild verifies that // Close blocks until every in-flight patched rebuild has released its // reservation, and that the rebuilder becomes terminal once all rebuilds // complete. It exercises both continuation and schema directive kinds with // multiple concurrent rebuilds to exercise the in-flight counter and condition // variable. func TestOpenAIRequestRebuilderConcurrentPatchedRebuildCloseCleanup(t *testing.T) { body := []byte(`{"model":"alias","messages":[{"role":"user","content":"hi"}]}`) maxBytes := int64(len(body) + 256) type testCase struct { name string directive func(ref streamgate.RecoveryRequestSnapshotRef) streamgate.RecoveryDirective patchCode func(store *openAIRecoveryPatchStore, ref streamgate.RecoveryRequestSnapshotRef) error } cases := []testCase{ { name: "continuation", directive: func(ref streamgate.RecoveryRequestSnapshotRef) streamgate.RecoveryDirective { d, _ := streamgate.NewRecoveryDirectiveContinuation(0, "snap.cont") return d }, patchCode: func(store *openAIRecoveryPatchStore, _ streamgate.RecoveryRequestSnapshotRef) error { return store.PutContinuation("snap.cont", 0, json.RawMessage(`[{"role":"assistant","content":"ok"}]`)) }, }, { name: "schema", directive: func(ref streamgate.RecoveryRequestSnapshotRef) streamgate.RecoveryDirective { d, _ := streamgate.NewRecoveryDirectiveSchema("snap.sch", "p.input") return d }, patchCode: func(store *openAIRecoveryPatchStore, _ streamgate.RecoveryRequestSnapshotRef) error { return store.PutSchema("snap.sch", "p.input", json.RawMessage(`[{"role":"user","content":"fixed"}]`)) }, }, } for _, tc := range cases { tc := tc t.Run(tc.name, func(t *testing.T) { _, rebuilder, ref := newOpenAIRebuilderFixture(t, openAIRebuildEndpointChat, body, maxBytes) if err := tc.patchCode(rebuilder.PatchStore(), ref); err != nil { t.Fatalf("patch store: %v", err) } directive := tc.directive(ref) var strategy streamgate.RecoveryStrategy if tc.name == "schema" { strategy = streamgate.RecoveryStrategySchemaRepair } else { strategy = streamgate.RecoveryStrategyContinuationRepair } plan := mustOpenAIRecoveryPlan(t, "plan."+tc.name, strategy, directive) // Start multiple rebuilds concurrently to exercise the in-flight // counter. Each rebuild takes a patch and completes quickly with a // non-cancelled context, but Close must wait for all of them. const nRebuilds = 3 var wg sync.WaitGroup wg.Add(nRebuilds) for i := 0; i < nRebuilds; i++ { go func() { defer wg.Done() rebuilder.RebuildRequest(context.Background(), ref, plan) }() } // Wait for all rebuilds to complete. wg.Wait() // Now call Close. It should return quickly since no rebuilds are // in-flight. rebuilder.Close() // After Close: store is closed and empty, reservation is zero. store := rebuilder.PatchStore() store.mu.Lock() storeClosed := store.closed storeCont := len(store.continuation) storeSche := len(store.schema) store.mu.Unlock() if !storeClosed { t.Fatal("patch store not closed after Close") } if storeCont != 0 || storeSche != 0 { t.Fatalf("patch store not empty after Close: cont=%d schema=%d", storeCont, storeSche) } rebuilt := rebuilder.RebuiltStore() rebuilt.mu.Lock() rebuiltLeases := len(rebuilt.leases) rebuilt.mu.Unlock() if rebuiltLeases != 0 { t.Fatalf("rebuilt store retained %d leases after Close", rebuiltLeases) } // Post-close rebuild must return ErrIngressSnapshotClosed. if _, err := rebuilder.RebuildRequest(context.Background(), ref, plan); !errors.Is(err, streamgate.ErrIngressSnapshotClosed) { t.Fatalf("post-close RebuildRequest = %v, want ErrIngressSnapshotClosed", err) } }) } } // TestOpenAIRequestRebuilderCloseBlocksForInFlightPatchedRebuild verifies that // Close blocks until an in-flight patched rebuild completes, using a context // that is cancelled after a short delay to force the rebuild to return with // context.Canceled. func TestOpenAIRequestRebuilderCloseBlocksForInFlightPatchedRebuild(t *testing.T) { body := []byte(`{"model":"alias","messages":[{"role":"user","content":"hi"}]}`) maxBytes := int64(len(body) + 256) // Test continuation directive _, rebuilder, ref := newOpenAIRebuilderFixture(t, openAIRebuildEndpointChat, body, maxBytes) patch := json.RawMessage(`[{"role":"assistant","content":"ok"}]`) if err := rebuilder.PatchStore().PutContinuation("snap.cont", 0, patch); err != nil { t.Fatalf("PutContinuation: %v", err) } directive, _ := streamgate.NewRecoveryDirectiveContinuation(0, "snap.cont") plan := mustOpenAIRecoveryPlan(t, "plan.block", streamgate.RecoveryStrategyContinuationRepair, directive) // Create a context that will be cancelled after a short delay. ctx, cancel := context.WithCancel(context.Background()) // Start a rebuild that will be interrupted by context cancellation. var wg sync.WaitGroup wg.Add(1) var buildErr error go func() { defer wg.Done() _, buildErr = rebuilder.RebuildRequest(ctx, ref, plan) }() // Give the rebuild time to start and take the patch. time.Sleep(50 * time.Millisecond) // Cancel the context to force the rebuild to return. cancel() // Wait for the rebuild to complete. wg.Wait() // Now call Close. It should return quickly since the rebuild has completed. rebuilder.Close() // Verify the rebuild returned context.Canceled (or ErrIngressSnapshotClosed // if Close ran first). if buildErr != nil && !errors.Is(buildErr, context.Canceled) && !errors.Is(buildErr, streamgate.ErrIngressSnapshotClosed) { t.Fatalf("build error = %v, want context.Canceled or ErrIngressSnapshotClosed", buildErr) } // After Close: store is closed and empty. store := rebuilder.PatchStore() store.mu.Lock() storeClosed := store.closed storeCont := len(store.continuation) storeSche := len(store.schema) store.mu.Unlock() if !storeClosed { t.Fatal("patch store not closed after Close") } if storeCont != 0 || storeSche != 0 { t.Fatalf("patch store not empty after Close: cont=%d schema=%d", storeCont, storeSche) } } // TestOpenAIRequestRebuilderNilReceiver verifies that nil receivers return // ErrIngressSnapshotClosed instead of panicking, covering the regression // introduced when mutex was added to the rebuilder. func TestOpenAIRequestRebuilderNilReceiverRegression(t *testing.T) { var rebuilder *openAIRequestRebuilder rebuilder.Close() ref, err := streamgate.NewRecoveryRequestSnapshotRef("nil.test", 0, 0, 1024) if err != nil { t.Fatalf("NewRecoveryRequestSnapshotRef: %v", err) } directive, _ := streamgate.NewRecoveryDirectiveExact(ref.SnapshotRef()) plan := mustOpenAIRecoveryPlan(t, "plan.nil", streamgate.RecoveryStrategyExactReplay, directive) if _, err := rebuilder.RebuildRequest(context.Background(), ref, plan); !errors.Is(err, streamgate.ErrIngressSnapshotClosed) { t.Fatalf("nil RebuildRequest = %v, want ErrIngressSnapshotClosed", err) } }