iop/apps/edge/internal/openai/stream_gate_vertical_slice_test.go
toki c380892159 fix(openai): 채팅 터널 실패 종료를 보장한다
재시도 소진 뒤 거부된 터미널을 재생하지 않고 SSE 오류로 종료해야 클라이언트가 중단된 스트림 대신 명시적인 실패를 관측할 수 있다.
2026-07-30 17:42:25 +09:00

5879 lines
224 KiB
Go

package openai
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"reflect"
"sort"
"strconv"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
"go.uber.org/zap"
"go.uber.org/zap/zapcore"
edgeservice "iop/apps/edge/internal/service"
"iop/packages/go/config"
"iop/packages/go/streamgate"
iop "iop/proto/gen/iop"
)
// --- test-only InjectedViolationFilter --------------------------------------
//
// injectedViolationFilter is the test-only filter named by the vertical-slice
// plan (SDD S16). It is deliberately kept local to this test file: production
// wiring (openAIStreamGateRegistrations) never registers it. It triggers an
// exact_replay RecoveryIntent on the first `trigger` text_delta evaluations it
// sees, then passes for the rest of the request.
//
// It uses terminal_gate hold mode rather than none: none mode is reserved
// for provider-error-only observers (its trigger-ready rule only fires on
// provider_error/terminal events regardless of content subscription — see
// packages/go/streamgate/evidence_tail.go deriveTriggerReady), so a filter
// that wants to judge accumulated content must hold it until the terminal
// trigger like a real output-validation filter would.
// injectedViolationPriority is shared by every test registration of
// injectedViolationFilter and its emitted RecoveryIntent: the Arbiter
// requires an exact match between a violation's intent priority and the
// filter's effective registration priority.
const injectedViolationPriority = 10
type injectedViolationFilter struct {
streamgate.FilterBase
requestRef string
continuationCursor int
passBeforeTrigger int
mu sync.Mutex
trigger int
}
func newInjectedViolationFilter(t *testing.T, id string, trigger int, requestRef string) *injectedViolationFilter {
t.Helper()
base, err := streamgate.NewFilterBase(id)
if err != nil {
t.Fatalf("NewFilterBase: %v", err)
}
return &injectedViolationFilter{FilterBase: base, requestRef: requestRef, trigger: trigger}
}
// newInjectedContinuationViolationFilter is the test-only stand-in for the
// follow-up repeat guard. It exercises S20's Rebuilder seam without making the
// semantic detector itself part of this task: it requests continuation from the
// request-local recovery source through the stable ingress snapshot reference.
func newInjectedContinuationViolationFilter(t *testing.T, id string, trigger int, requestRef string, cursor int) *injectedViolationFilter {
t.Helper()
filter := newInjectedViolationFilter(t, id, trigger, requestRef)
filter.continuationCursor = cursor
// A rolling filter needs one safe evaluation to open the downstream stream
// before a continuation plan can use the Core's continue-stream mode.
filter.passBeforeTrigger = 1
return filter
}
func (f *injectedViolationFilter) Applies(streamgate.FilterContext) bool { return true }
func (f *injectedViolationFilter) HoldRequirement(streamgate.FilterContext) streamgate.FilterHoldRequirement {
if f.continuationCursor > 0 {
req, err := streamgate.NewFilterHoldRequirementRolling(streamGateChannelDefault, []streamgate.EventKind{streamgate.EventKindTextDelta}, 1)
if err != nil {
panic(err)
}
return req
}
req, err := streamgate.NewFilterHoldRequirementTerminalGate(streamGateChannelDefault, []streamgate.EventKind{streamgate.EventKindTextDelta}, streamgate.EventKindTerminal)
if err != nil {
panic(err)
}
return req
}
func (f *injectedViolationFilter) Evaluate(ctx context.Context, fctx streamgate.FilterContext, batch streamgate.EvidenceBatch) (streamgate.FilterDecision, error) {
f.mu.Lock()
fire := f.passBeforeTrigger == 0 && f.trigger > 0
if f.passBeforeTrigger > 0 {
f.passBeforeTrigger--
} else if fire {
f.trigger--
}
f.mu.Unlock()
events := batch.Events()
kind := streamgate.EventKindTextDelta
if len(events) > 0 {
kind = events[len(events)-1].Kind()
}
evidence, err := streamgate.NewSanitizedEvidence(kind, streamGateChannelDefault, "test.rule", "test_injected_violation", streamgate.FixedFingerprint{0x02}, 1, 0, streamgate.FilterOutcomeKindEvaluated, batch.CapturedAt())
if err != nil {
return streamgate.FilterDecision{}, err
}
if !fire {
return streamgate.NewFilterDecision(streamgate.FilterDecisionKindPass, "test.consumer", f.ID(), "test.rule", evidence, nil)
}
var directive streamgate.RecoveryDirective
if f.continuationCursor > 0 {
directive, err = streamgate.NewRecoveryDirectiveContinuation(f.continuationCursor, f.requestRef)
} else {
directive, err = streamgate.NewRecoveryDirectiveExact(f.requestRef)
}
if err != nil {
return streamgate.FilterDecision{}, err
}
strategy := streamgate.RecoveryStrategyExactReplay
if f.continuationCursor > 0 {
strategy = streamgate.RecoveryStrategyContinuationRepair
}
intent, err := streamgate.NewRecoveryIntent(strategy, directive, "test_injected", injectedViolationPriority)
if err != nil {
return streamgate.FilterDecision{}, err
}
return streamgate.NewFilterDecision(streamgate.FilterDecisionKindViolation, "test.consumer", f.ID(), "test.rule", evidence, &intent)
}
var _ streamgate.Filter = (*injectedViolationFilter)(nil)
func streamGateTestRegistry(t *testing.T, extra ...streamgate.FilterRegistration) streamgate.FilterRegistrySnapshot {
t.Helper()
regs, err := openAIStreamGateNoopRegistrations()
if err != nil {
t.Fatalf("openAIStreamGateNoopRegistrations: %v", err)
}
regs = append(regs, extra...)
snap, err := streamgate.NewFilterRegistrySnapshot(streamGateConfigGeneration, regs, nil)
if err != nil {
t.Fatalf("NewFilterRegistrySnapshot: %v", err)
}
return snap
}
// --- recording ResponseWriter -------------------------------------------------
// recordingResponseWriter records the ordered sequence of Header/WriteHeader/
// Write calls so tests can assert that no status/header/role write happens
// before the Core's first safe release.
type recordingResponseWriter struct {
header http.Header
body bytes.Buffer
code int
calls []string
}
func newRecordingResponseWriter() *recordingResponseWriter {
return &recordingResponseWriter{header: make(http.Header)}
}
func (w *recordingResponseWriter) Header() http.Header { return w.header }
func (w *recordingResponseWriter) WriteHeader(code int) {
w.code = code
w.calls = append(w.calls, fmt.Sprintf("header:%d", code))
}
func (w *recordingResponseWriter) Write(p []byte) (int, error) {
w.calls = append(w.calls, fmt.Sprintf("write:%d", len(p)))
return w.body.Write(p)
}
func (w *recordingResponseWriter) Flush() {}
func (w *recordingResponseWriter) headerCallCount() int {
n := 0
for _, c := range w.calls {
if strings.HasPrefix(c, "header:") {
n++
}
}
return n
}
type notifyingResponseWriter struct {
*recordingResponseWriter
writes chan struct{}
}
func newNotifyingResponseWriter() *notifyingResponseWriter {
return &notifyingResponseWriter{recordingResponseWriter: newRecordingResponseWriter(), writes: make(chan struct{}, 8)}
}
func (w *notifyingResponseWriter) Write(p []byte) (int, error) {
w.calls = append(w.calls, fmt.Sprintf("write:%d", len(p)))
n, err := w.body.Write(p)
select {
case w.writes <- struct{}{}:
default:
}
return n, err
}
var (
_ http.ResponseWriter = (*recordingResponseWriter)(nil)
_ http.Flusher = (*recordingResponseWriter)(nil)
)
// --- SSE parsing helper -------------------------------------------------------
func parseSSEChatChunks(t *testing.T, body string) []chatCompletionChunk {
t.Helper()
var chunks []chatCompletionChunk
for _, frame := range strings.Split(body, "\n\n") {
frame = strings.TrimSpace(frame)
if frame == "" {
continue
}
payload := strings.TrimPrefix(frame, "data: ")
if payload == "[DONE]" {
continue
}
var chunk chatCompletionChunk
if err := json.Unmarshal([]byte(payload), &chunk); err != nil {
continue // error chunks decode into errorResponse, not chatCompletionChunk; skip
}
if len(chunk.Choices) == 0 {
continue
}
chunks = append(chunks, chunk)
}
return chunks
}
// parseCompleteResponsesSSE rejects any byte that is not part of a complete
// Responses SSE frame. Unlike the chat helper above, it does not forgive a
// malformed payload: recovery tests use it to guard the caller wire boundary.
func parseCompleteResponsesSSE(t *testing.T, body string) []map[string]any {
t.Helper()
remaining := []byte(body)
var payloads []map[string]any
for len(remaining) > 0 {
frame, rest, ok := takeOpenAISSEFrame(remaining)
if !ok {
t.Fatalf("Responses SSE has a non-frame remainder: %q", remaining)
}
if !bytes.HasPrefix(frame, []byte("data: ")) {
t.Fatalf("Responses SSE frame does not begin with data: %q", frame)
}
data := openAISSEData(frame)
if data != "[DONE]" {
var payload map[string]any
if err := json.Unmarshal([]byte(data), &payload); err != nil {
t.Fatalf("Responses SSE payload is not JSON: %q: %v", data, err)
}
payloads = append(payloads, payload)
}
remaining = rest
}
return payloads
}
// responsesOracleItem records the event-specific state of one output item so
// the lifecycle oracle can enforce expected type/status, stable identity and
// indexes, legal transition order, exact done counts, and value agreement
// between every accumulated delta, its done payload, the completed item, and
// the terminal output.
type responsesOracleItem struct {
id string
itemType string
outputIndex int
contentIndex int
callID string
name string
role string
contentOpened bool
value string
textDone int
contentDone int
argumentsDone int
itemDone int
}
// assertResponsesSSELifecycle validates the full public Responses SSE contract
// for a caller stream and returns the parsed payloads. wantTerminal is the
// required last event type ("response.completed" or "error").
func assertResponsesSSELifecycle(t *testing.T, body, wantTerminal string) []map[string]any {
t.Helper()
payloads := parseCompleteResponsesSSE(t, body)
if got := strings.Count(body, "data: [DONE]\n\n"); got != 1 {
t.Fatalf("Responses [DONE] count = %d, want 1; body=%q", got, body)
}
var previousSequence float64
terminalCount := 0
sawCompleted := false
responseID := ""
items := make(map[string]*responsesOracleItem)
order := make([]string, 0)
indexOwner := make(map[int]string)
for _, payload := range payloads {
kind, _ := payload["type"].(string)
sequence, ok := payload["sequence_number"].(float64)
if !ok || sequence <= previousSequence {
t.Fatalf("Responses event %q sequence_number = %#v after %v; payload=%#v", kind, payload["sequence_number"], previousSequence, payload)
}
previousSequence = sequence
requireString := func(key string) {
t.Helper()
if value, ok := payload[key].(string); !ok || value == "" {
t.Fatalf("Responses event %q missing %s: %#v", kind, key, payload)
}
}
requireNumber := func(key string) {
t.Helper()
if _, ok := payload[key].(float64); !ok {
t.Fatalf("Responses event %q missing %s: %#v", kind, key, payload)
}
}
forbid := func(key string) {
t.Helper()
if _, ok := payload[key]; ok {
t.Fatalf("Responses event %q must not carry %s: %#v", kind, key, payload)
}
}
known := func(itemID string) *responsesOracleItem {
t.Helper()
it, ok := items[itemID]
if !ok {
t.Fatalf("Responses event %q references unknown item %q: %#v", kind, itemID, payload)
}
return it
}
assertOutputIndex := func(it *responsesOracleItem) {
t.Helper()
if got := int(payload["output_index"].(float64)); got != it.outputIndex {
t.Fatalf("Responses event %q output_index = %d, want stable %d for item %q: %#v", kind, got, it.outputIndex, it.id, payload)
}
}
assertContentIndex := func(it *responsesOracleItem) {
t.Helper()
if err := validateResponsesContentIndex(payload, it.contentIndex); err != nil {
t.Fatalf("Responses event %q %v", kind, err)
}
}
switch kind {
case "response.created":
response, ok := payload["response"].(map[string]any)
if !ok || response["id"] == "" {
t.Fatalf("response.created has no response identity: %#v", payload)
}
id, _ := response["id"].(string)
if responseID != "" && responseID != id {
t.Fatalf("response.created identity = %q, want stable %q", id, responseID)
}
responseID = id
case "response.output_item.added":
requireNumber("output_index")
item, ok := payload["item"].(map[string]any)
id, _ := item["id"].(string)
itemType, _ := item["type"].(string)
status, _ := item["status"].(string)
if !ok || id == "" || itemType == "" {
t.Fatalf("response.output_item.added has no item identity: %#v", payload)
}
if status != "in_progress" {
t.Fatalf("response.output_item.added status = %q, want in_progress: %#v", status, payload)
}
if _, exists := items[id]; exists {
t.Fatalf("response.output_item.added duplicates item %q", id)
}
index := int(payload["output_index"].(float64))
if prior, exists := indexOwner[index]; exists {
t.Fatalf("response.output_item.added reuses output index %d for %q and %q", index, prior, id)
}
rec := &responsesOracleItem{id: id, itemType: itemType, outputIndex: index}
switch itemType {
case "function_call":
callID, _ := item["call_id"].(string)
name, _ := item["name"].(string)
if callID == "" || name == "" {
t.Fatalf("function output_item.added missing call_id/name: %#v", item)
}
rec.callID = callID
rec.name = name
case "message":
role, _ := item["role"].(string)
if role == "" {
t.Fatalf("message output_item.added missing role: %#v", item)
}
rec.role = role
if content, ok := item["content"].([]any); !ok || len(content) != 0 {
t.Fatalf("message output_item.added opening content must be empty: %#v", item)
}
case "reasoning":
if content, ok := item["content"].([]any); !ok || len(content) != 0 {
t.Fatalf("reasoning output_item.added opening content must be empty: %#v", item)
}
default:
t.Fatalf("response.output_item.added has an unsupported item type %q: %#v", itemType, item)
}
items[id] = rec
order = append(order, id)
indexOwner[index] = id
case "response.content_part.added":
requireString("item_id")
requireNumber("output_index")
requireNumber("content_index")
it := known(payload["item_id"].(string))
assertOutputIndex(it)
if it.itemType == "function_call" {
t.Fatalf("response.content_part.added opened a function item: %#v", payload)
}
assertResponsesPartShape(t, it.itemType, payload["part"], "")
it.contentOpened = true
idx, err := parseResponsesContentIndex(payload)
if err != nil {
t.Fatalf("response.content_part.added invalid content_index: %v", err)
}
it.contentIndex = idx
case "response.output_text.delta", "response.reasoning_text.delta":
requireString("item_id")
requireNumber("output_index")
requireString("delta")
it := known(payload["item_id"].(string))
assertOutputIndex(it)
assertContentIndex(it)
assertResponsesTextEventKind(t, kind, it, "delta")
if !it.contentOpened || it.textDone > 0 {
t.Fatalf("Responses text delta has no opened content or follows a done: %#v", payload)
}
it.value += payload["delta"].(string)
case "response.function_call_arguments.delta":
requireString("item_id")
requireNumber("output_index")
requireString("delta")
forbid("call_id")
forbid("name")
it := known(payload["item_id"].(string))
assertOutputIndex(it)
if it.itemType != "function_call" || it.argumentsDone > 0 {
t.Fatalf("function delta references a non-function or completed item: %#v", payload)
}
it.value += payload["delta"].(string)
case "response.output_text.done", "response.reasoning_text.done":
requireString("item_id")
requireNumber("output_index")
requireString("text")
it := known(payload["item_id"].(string))
assertOutputIndex(it)
assertContentIndex(it)
assertResponsesTextEventKind(t, kind, it, "done")
if !it.contentOpened {
t.Fatalf("Responses text done has no opened content: %#v", payload)
}
if got := payload["text"]; got != it.value {
t.Fatalf("Responses text done = %#v, want accumulated delta %q: %#v", got, it.value, payload)
}
it.textDone++
case "response.content_part.done":
requireString("item_id")
requireNumber("output_index")
it := known(payload["item_id"].(string))
assertOutputIndex(it)
assertContentIndex(it)
if !it.contentOpened || it.textDone == 0 || it.contentDone > 0 {
t.Fatalf("Responses content done is out of order: %#v", payload)
}
assertResponsesPartShape(t, it.itemType, payload["part"], it.value)
it.contentDone++
case "response.function_call_arguments.done":
requireString("item_id")
requireNumber("output_index")
requireString("name")
requireString("arguments")
it := known(payload["item_id"].(string))
assertOutputIndex(it)
if it.itemType != "function_call" {
t.Fatalf("function arguments done references a non-function item: %#v", payload)
}
if got := payload["arguments"]; got != it.value {
t.Fatalf("function arguments done = %#v, want accumulated delta %q: %#v", got, it.value, payload)
}
if got := payload["name"]; got != it.name {
t.Fatalf("function arguments done name = %#v, want %q: %#v", got, it.name, payload)
}
it.argumentsDone++
case "response.output_item.done":
requireNumber("output_index")
item, ok := payload["item"].(map[string]any)
id, _ := item["id"].(string)
if !ok || id == "" {
t.Fatalf("response.output_item.done has no item identity: %#v", payload)
}
it := known(id)
assertOutputIndex(it)
if status, _ := item["status"].(string); status != "completed" {
t.Fatalf("response.output_item.done status = %#v, want completed: %#v", item["status"], payload)
}
if it.itemDone > 0 {
t.Fatalf("response.output_item.done repeats item %q", id)
}
if it.itemType == "function_call" {
if it.argumentsDone == 0 {
t.Fatalf("function output_item.done precedes arguments done: %#v", payload)
}
} else if it.textDone == 0 || it.contentDone == 0 {
t.Fatalf("output_item.done precedes text/content done: %#v", payload)
}
assertResponsesCompletedItem(t, it, item)
it.itemDone++
case "response.completed":
terminalCount++
sawCompleted = true
response, ok := payload["response"].(map[string]any)
if !ok || response["id"] == "" || response["object"] != "response" || response["status"] != "completed" || response["output"] == nil || response["usage"] == nil {
t.Fatalf("response.completed has an incomplete response object: %#v", payload)
}
if responseID != "" && response["id"] != responseID {
t.Fatalf("response.completed identity = %q, want %q", response["id"], responseID)
}
for _, id := range order {
if items[id].itemDone != 1 {
t.Fatalf("response.completed preceded item completion for %q: %#v", id, payloads)
}
}
output, ok := response["output"].([]any)
if !ok {
t.Fatalf("response.completed output is not an array: %#v", response)
}
expectedOrder := append([]string(nil), order...)
sort.Slice(expectedOrder, func(i, j int) bool {
return items[expectedOrder[i]].outputIndex < items[expectedOrder[j]].outputIndex
})
if len(output) != len(expectedOrder) {
t.Fatalf("response.completed output length = %d, want %d: %#v", len(output), len(expectedOrder), output)
}
for position, rawItem := range output {
item, ok := rawItem.(map[string]any)
if !ok {
t.Fatalf("response.completed output item is invalid: %#v", rawItem)
}
id, _ := item["id"].(string)
if id != expectedOrder[position] {
t.Fatalf("response.completed output item %d identity = %q, want %q: %#v", position, id, expectedOrder[position], item)
}
it := known(id)
assertResponsesCompletedItem(t, it, item)
}
case "error":
terminalCount++
requireString("code")
requireString("message")
default:
t.Fatalf("unexpected Responses SSE event type %q: %#v", kind, payload)
}
}
if terminalCount != 1 {
t.Fatalf("Responses terminal event count = %d, want 1; payloads=%#v", terminalCount, payloads)
}
if payloads[len(payloads)-1]["type"] != wantTerminal {
t.Fatalf("Responses terminal type = %q, want %q; payloads=%#v", payloads[len(payloads)-1]["type"], wantTerminal, payloads)
}
if sawCompleted {
for _, id := range order {
it := items[id]
if it.itemType == "function_call" {
if it.argumentsDone != 1 || it.itemDone != 1 {
t.Fatalf("function item %q done counts = (args %d, item %d), want exactly one each", id, it.argumentsDone, it.itemDone)
}
continue
}
if it.textDone != 1 || it.contentDone != 1 || it.itemDone != 1 {
t.Fatalf("item %q done counts = (text %d, content %d, item %d), want exactly one each", id, it.textDone, it.contentDone, it.itemDone)
}
}
}
return payloads
}
// assertResponsesTextEventKind rejects a text lifecycle event whose type does
// not match the item's type (output_text for message, reasoning_text for
// reasoning).
func assertResponsesTextEventKind(t *testing.T, kind string, it *responsesOracleItem, phase string) {
t.Helper()
want := "response.output_text." + phase
if it.itemType == "reasoning" {
want = "response.reasoning_text." + phase
}
if kind != want {
t.Fatalf("Responses %s event %q does not match item %q of type %q", phase, kind, it.id, it.itemType)
}
}
// assertResponsesPartShape validates a content part against the item type. An
// output-text part must carry annotations/logprobs arrays; a reasoning part
// must not. wantText is compared when non-empty is expected (done parts).
func assertResponsesPartShape(t *testing.T, itemType string, raw any, wantText string) {
t.Helper()
part, ok := raw.(map[string]any)
if !ok {
t.Fatalf("Responses content part is not an object: %#v", raw)
}
wantType := "output_text"
if itemType == "reasoning" {
wantType = "reasoning_text"
}
if part["type"] != wantType {
t.Fatalf("Responses content part type = %#v, want %q: %#v", part["type"], wantType, part)
}
if got, ok := part["text"].(string); !ok || got != wantText {
t.Fatalf("Responses content part text = %#v, want %q: %#v", part["text"], wantText, part)
}
if itemType == "reasoning" {
return
}
if _, ok := part["annotations"].([]any); !ok {
t.Fatalf("output_text content part missing annotations array: %#v", part)
}
if _, ok := part["logprobs"].([]any); !ok {
t.Fatalf("output_text content part missing logprobs array: %#v", part)
}
}
// assertResponsesCompletedItem validates a completed item object (in
// response.output_item.done and terminal output) against the accumulated
// oracle state.
func assertResponsesCompletedItem(t *testing.T, it *responsesOracleItem, item map[string]any) {
t.Helper()
if item["id"] != it.id || item["type"] != it.itemType || item["status"] != "completed" {
t.Fatalf("completed item identity/type/status = (%#v,%#v,%#v), want (%q,%q,%q): %#v", item["id"], item["type"], item["status"], it.id, it.itemType, "completed", item)
}
if it.itemType == "function_call" {
if item["arguments"] != it.value {
t.Fatalf("completed function arguments = %#v, want %q: %#v", item["arguments"], it.value, item)
}
if item["call_id"] != it.callID || item["name"] != it.name {
t.Fatalf("completed function metadata = (%#v,%#v), want (%q,%q): %#v", item["call_id"], item["name"], it.callID, it.name, item)
}
return
}
content, ok := item["content"].([]any)
if !ok || len(content) != 1 {
t.Fatalf("completed content item is invalid: %#v", item)
}
assertResponsesPartShape(t, it.itemType, content[0], it.value)
}
// parseResponsesContentIndex checks that payload["content_index"] is present,
// numeric (float64 in parsed JSON payload), and an exact integer. It returns
// the integer value or a descriptive error.
func parseResponsesContentIndex(payload map[string]any) (int, error) {
raw, ok := payload["content_index"]
if !ok {
return 0, fmt.Errorf("missing content_index: %#v", payload)
}
floatVal, ok := raw.(float64)
if !ok {
return 0, fmt.Errorf("non-numeric content_index type %T (%#v): %#v", raw, raw, payload)
}
if floatVal != float64(int(floatVal)) {
return 0, fmt.Errorf("non-integer content_index value %v: %#v", floatVal, payload)
}
return int(floatVal), nil
}
// validateResponsesContentIndex checks that payload["content_index"] is present,
// numeric (float64 in parsed JSON payload), an exact integer, and equals expectedIndex.
// It returns a descriptive error for missing, non-numeric, non-integer, or changed values.
func validateResponsesContentIndex(payload map[string]any, expectedIndex int) error {
got, err := parseResponsesContentIndex(payload)
if err != nil {
return err
}
if got != expectedIndex {
return fmt.Errorf("content_index = %d, want stable %d: %#v", got, expectedIndex, payload)
}
return nil
}
func TestResponsesLifecycleContentIndexInvariant(t *testing.T) {
tests := []struct {
name string
payload map[string]any
expectedIndex int
wantErr bool
}{
{
name: "matching content_index",
payload: map[string]any{
"content_index": float64(0),
},
expectedIndex: 0,
wantErr: false,
},
{
name: "matching content_index non-zero",
payload: map[string]any{
"content_index": float64(2),
},
expectedIndex: 2,
wantErr: false,
},
{
name: "changed content_index",
payload: map[string]any{
"content_index": float64(9),
},
expectedIndex: 0,
wantErr: true,
},
{
name: "missing content_index",
payload: map[string]any{
"item_id": "msg_1",
},
expectedIndex: 0,
wantErr: true,
},
{
name: "non-numeric string content_index",
payload: map[string]any{
"content_index": "0",
},
expectedIndex: 0,
wantErr: true,
},
{
name: "non-numeric nil content_index",
payload: map[string]any{
"content_index": nil,
},
expectedIndex: 0,
wantErr: true,
},
{
name: "fractional content_index matching truncated integer",
payload: map[string]any{
"content_index": float64(0.5),
},
expectedIndex: 0,
wantErr: true,
},
{
name: "fractional content_index non-zero",
payload: map[string]any{
"content_index": float64(2.5),
},
expectedIndex: 2,
wantErr: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
err := validateResponsesContentIndex(tt.payload, tt.expectedIndex)
if (err != nil) != tt.wantErr {
t.Errorf("validateResponsesContentIndex() error = %v, wantErr %v", err, tt.wantErr)
}
})
}
}
func responsesSSEOutputText(payloads []map[string]any) string {
var output strings.Builder
for _, payload := range payloads {
if payload["type"] != "response.output_text.delta" {
continue
}
text, _ := payload["delta"].(string)
output.WriteString(text)
}
return output.String()
}
func joinedContent(chunks []chatCompletionChunk) string {
var b strings.Builder
for _, c := range chunks {
if len(c.Choices) > 0 {
b.WriteString(c.Choices[0].Delta.Content)
}
}
return b.String()
}
// --- chat handler-level: disabled vs enabled Noop parity ---------------------
func streamGateChatHandlerFixture(t *testing.T, enabled bool) (*Server, *fakeRunService) {
t.Helper()
fake := &fakeRunService{events: make(chan *iop.RunEvent, 3)}
fake.events <- &iop.RunEvent{Type: "delta", Delta: "hello"}
fake.events <- &iop.RunEvent{Type: "delta", Delta: " world"}
fake.events <- &iop.RunEvent{Type: "complete", Usage: &iop.Usage{InputTokens: 1, OutputTokens: 2}}
srv := NewServer(config.EdgeOpenAIConf{
Adapter: "ollama",
Target: "llama-fixed",
SessionID: "cli",
TimeoutSec: 15,
StreamEvidenceGate: config.StreamEvidenceGateConf{
Enabled: enabled,
},
}, fake, nil)
return srv, fake
}
func TestStreamGateChatHandlerDisabledVsEnabledNoopPass(t *testing.T) {
body := `{"model":"client-model","messages":[{"role":"user","content":"hi"}],"stream":true}`
for _, enabled := range []bool{false, true} {
enabled := enabled
name := "disabled"
if enabled {
name = "enabled"
}
t.Run(name, func(t *testing.T) {
srv, fake := streamGateChatHandlerFixture(t, enabled)
req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(body))
w := httptest.NewRecorder()
srv.handleChatCompletions(w, req)
if w.Code != http.StatusOK {
t.Fatalf("status: got %d body=%s", w.Code, w.Body.String())
}
chunks := parseSSEChatChunks(t, w.Body.String())
if len(chunks) == 0 {
t.Fatalf("no SSE chunks decoded from body=%s", w.Body.String())
}
if chunks[0].Choices[0].Delta.Role != "assistant" {
t.Fatalf("first chunk role: got %+v", chunks[0].Choices[0].Delta)
}
if got := joinedContent(chunks); got != "hello world" {
t.Fatalf("joined content: got %q", got)
}
last := chunks[len(chunks)-1]
if last.Choices[0].FinishReason != "stop" {
t.Fatalf("finish reason: got %q", last.Choices[0].FinishReason)
}
if !strings.Contains(w.Body.String(), "data: [DONE]") {
t.Fatalf("missing [DONE] sentinel: body=%s", w.Body.String())
}
if len(fake.reqsSnapshot()) != 1 {
t.Fatalf("dispatch attempts: got %d, want 1", len(fake.reqsSnapshot()))
}
})
}
}
// --- adapter-level: chat runtime + injected violation recovery ---------------
// buildStreamGateChatFixture builds a ready-to-run chatDispatchContext plus the
// initial fakeRunResult a handler would already have dispatched, mirroring
// production ingress without going through the HTTP decode path.
func buildStreamGateChatFixture(t *testing.T, enabled bool, maxRequestFaultRecovery int) (*Server, *chatDispatchContext, *fakeRunResult, *fakeRunService) {
t.Helper()
rawBody := []byte(`{"model":"client-model","messages":[{"role":"user","content":"hi"}],"stream":true}`)
var req chatCompletionRequest
if err := json.Unmarshal(rawBody, &req); err != nil {
t.Fatalf("decode fixture body: %v", err)
}
fault := maxRequestFaultRecovery
fake := &fakeRunService{}
srv := NewServer(config.EdgeOpenAIConf{
Adapter: "ollama",
Target: "llama-fixed",
SessionID: "cli",
TimeoutSec: 15,
StreamEvidenceGate: config.StreamEvidenceGateConf{
Enabled: enabled,
MaxRequestFaultRecovery: &fault,
},
}, fake, nil)
route := routeDispatch{Adapter: "ollama", Target: "llama-fixed", SessionID: "cli", TimeoutSec: 15}
requestCtx := newTestRequestContext(t, route, rawBody)
basePrompt := promptFromMessages(req.Messages)
outputPolicy := srv.resolveOutputPolicy(basePrompt)
dc := srv.newChatDispatchContext(requestCtx, req, basePrompt, outputPolicy)
initial := &fakeRunResult{
dispatch: edgeservice.RunDispatch{RunID: "run-attempt-1", Adapter: "ollama", Target: "llama-fixed", SessionID: "cli", TimeoutSec: 15},
events: bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "should-be-discarded"},
&iop.RunEvent{Type: "complete"},
),
}
return srv, dc, initial, fake
}
func TestStreamGateChatBuildRuntimeInjectedViolationSingleRetrySuccess(t *testing.T) {
srv, dc, initial, fake := buildStreamGateChatFixture(t, true, 3)
fake.eventRuns = []chan *iop.RunEvent{bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "final answer"},
&iop.RunEvent{Type: "complete"},
)}
fake.runIDs = []string{"run-attempt-2"}
snapRef, err := dc.ingress.recoveryRef()
if err != nil {
t.Fatalf("recoveryRef: %v", err)
}
violation := newInjectedViolationFilter(t, "test.injected_violation", 1, snapRef.SnapshotRef())
reg, err := streamgate.NewFilterRegistration(violation, streamGateNoopCapability, true, streamgate.FilterEnforcementBlocking, streamGateFilterTimeout, injectedViolationPriority)
if err != nil {
t.Fatalf("NewFilterRegistration: %v", err)
}
registry := streamGateTestRegistry(t, reg)
w := newRecordingResponseWriter()
rt, usage, err := srv.buildOpenAIChatStreamGateRuntime(dc, initial, newOpenAIChatSSEReleaseSink(w, w, "chatcmpl-run-attempt-1", time.Now().Unix(), "client-model"), registry)
if err != nil {
t.Fatalf("buildOpenAIChatStreamGateRuntime: %v", err)
}
if err := rt.Run(context.Background()); err != nil {
t.Fatalf("rt.Run: %v", err)
}
_ = usage
if got := len(fake.reqsSnapshot()); got != 1 {
t.Fatalf("recovery dispatch count: got %d, want 1 (exact_replay resubmits once)", got)
}
if got := w.headerCallCount(); got != 1 {
t.Fatalf("WriteHeader call count: got %d, want 1 (no eager header from the discarded first attempt)", got)
}
chunks := parseSSEChatChunks(t, w.body.String())
roleCount := 0
for _, c := range chunks {
if c.Choices[0].Delta.Role == "assistant" {
roleCount++
}
}
if roleCount != 1 {
t.Fatalf("assistant role chunk count: got %d, want 1", roleCount)
}
if got := joinedContent(chunks); got != "final answer" {
t.Fatalf("joined content: got %q, want only the retry attempt's content (no leak from the discarded first attempt)", got)
}
}
func TestOpenAIRepeatResumeBuildRunsAfterAbort(t *testing.T) {
srv, dc, initial, fake := buildStreamGateChatFixture(t, true, 3)
recorder := &recordingOpenAIObservationSink{}
srv.SetObservationSink(recorder)
// The catalog is the request-start source of the context window used by the
// resume builder. The direct fixture still uses the normalized fake service.
srv.SetModelCatalog([]config.ModelCatalogEntry{{ID: "client-model", ContextWindowTokens: 4096}})
initial.events = bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "safe prefix "},
&iop.RunEvent{Type: "delta", Delta: "repeated-tail"},
&iop.RunEvent{Type: "complete"},
)
fake.eventRuns = []chan *iop.RunEvent{bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "recovered answer"},
&iop.RunEvent{Type: "complete"},
)}
fake.runIDs = []string{"run-attempt-2"}
snapRef, err := dc.ingress.recoveryRef()
if err != nil {
t.Fatalf("recoveryRef: %v", err)
}
violation := newInjectedContinuationViolationFilter(t, "test.resume", 1, snapRef.SnapshotRef(), len("safe prefix "))
reg, err := streamgate.NewFilterRegistration(violation, streamGateNoopCapability, true, streamgate.FilterEnforcementBlocking, streamGateFilterTimeout, injectedViolationPriority)
if err != nil {
t.Fatalf("NewFilterRegistration: %v", err)
}
w := newRecordingResponseWriter()
rt, _, err := srv.buildOpenAIChatStreamGateRuntime(dc, initial, newOpenAIChatSSEReleaseSink(w, w, "chatcmpl-run-attempt-1", time.Now().Unix(), "client-model"), streamGateTestRegistry(t, reg))
if err != nil {
t.Fatalf("buildOpenAIChatStreamGateRuntime: %v", err)
}
if err := rt.Run(context.Background()); err != nil {
t.Fatalf("rt.Run: %v", err)
}
if got := len(fake.reqsSnapshot()); got != 1 {
t.Fatalf("recovery dispatch count = %d, want 1 after the aborted attempt; response=%s", got, w.body.String())
}
if got := joinedContent(parseSSEChatChunks(t, w.body.String())); !strings.Contains(got, "recovered answer") {
t.Fatalf("released content = %q, want the continuation attempt after abort", got)
}
var recovery []streamgate.ObservationKind
for _, observation := range recorder.Snapshot() {
switch observation.Kind() {
case streamgate.ObservationKindRecoveryAttemptAborted,
streamgate.ObservationKindRecoveryRebuilt,
streamgate.ObservationKindRecoveryDispatched:
recovery = append(recovery, observation.Kind())
}
}
wantRecovery := []streamgate.ObservationKind{
streamgate.ObservationKindRecoveryAttemptAborted,
streamgate.ObservationKindRecoveryRebuilt,
streamgate.ObservationKindRecoveryDispatched,
}
if len(recovery) != len(wantRecovery) {
t.Fatalf("recovery lifecycle rows = %v, want %v", recovery, wantRecovery)
}
for i := range wantRecovery {
if recovery[i] != wantRecovery[i] {
t.Fatalf("recovery lifecycle order = %v, want %v", recovery, wantRecovery)
}
}
}
func TestRepeatGuardStreamOpenNoDuplicatePrefix(t *testing.T) {
srv, dc, initial, fake := buildStreamGateChatFixture(t, true, 3)
srv.SetModelCatalog([]config.ModelCatalogEntry{{ID: "client-model", ContextWindowTokens: 8192}})
safePrefix := uniqueKoreanRunes(520)
initial.events = bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: safePrefix},
&iop.RunEvent{Type: "delta", Delta: safePrefix},
&iop.RunEvent{Type: "complete"},
)
fake.eventRuns = []chan *iop.RunEvent{bufferedRunEvents(
// The replacement provider echoes the assistant prefill. The
// request-local source must suppress that known prefix once.
&iop.RunEvent{Type: "delta", Delta: safePrefix + "recovered suffix"},
&iop.RunEvent{Type: "complete"},
)}
fake.runIDs = []string{"run-repeat-recovery"}
snapRef, err := dc.ingress.recoveryRef()
if err != nil {
t.Fatalf("recoveryRef: %v", err)
}
gateCfg := config.StreamEvidenceGateConf{
Enabled: true,
Filters: []config.StreamGateFilterPolicyConf{{
Filter: config.StreamGateFilterRepeatGuard,
Enforcement: config.StreamGateFilterEnforcementBlocking,
Priority: 10,
HoldEvidenceRunes: 500,
}},
}
registry, err := openAIStreamGateRegistrySnapshotFor(gateCfg, openAIOutputFilterContext{
endpoint: openAIRebuildEndpointChat,
requestRef: snapRef.SnapshotRef(),
})
if err != nil {
t.Fatalf("openAIStreamGateRegistrySnapshotFor: %v", err)
}
w := newRecordingResponseWriter()
rt, _, err := srv.buildOpenAIChatStreamGateRuntime(
dc,
initial,
newOpenAIChatSSEReleaseSink(w, w, "chatcmpl-repeat", time.Now().Unix(), "client-model"),
registry,
)
if err != nil {
t.Fatalf("buildOpenAIChatStreamGateRuntime: %v", err)
}
if err := rt.Run(context.Background()); err != nil {
t.Fatalf("rt.Run: %v", err)
}
if err := rt.CloseRequestResources(context.Background(), true); err != nil {
t.Fatalf("CloseRequestResources: %v", err)
}
requests := fake.reqsSnapshot()
if len(requests) != 1 {
t.Fatalf("recovery dispatches = %d, want 1", len(requests))
}
resumeMessages, _ := requests[0].Input["messages"].([]any)
if len(resumeMessages) < 1 {
t.Fatalf("recovery messages = %#v", requests[0].Input["messages"])
}
resumeAssistant, _ := resumeMessages[0].(map[string]any)
if got := resumeAssistant["content"]; got != safePrefix {
t.Fatalf("recovery safe prefix length/value = (%T, %d), want %d bytes", got, len(fmt.Sprint(got)), len(safePrefix))
}
options, _ := requests[0].Input["options"].(map[string]any)
if got := options["temperature"]; got != 0.2 {
t.Fatalf("recovery temperature = %#v, want 0.2", got)
}
body := w.body.String()
chunks := parseSSEChatChunks(t, body)
if got := joinedContent(chunks); got != safePrefix+"recovered suffix" {
t.Fatalf("downstream content length=%d, want safe prefix once plus recovered suffix", len(got))
}
roleCount := 0
for _, chunk := range chunks {
if len(chunk.Choices) > 0 && chunk.Choices[0].Delta.Role == "assistant" {
roleCount++
}
}
if roleCount != 1 {
t.Fatalf("assistant opening count = %d, want 1", roleCount)
}
if got := strings.Count(body, "data: [DONE]"); got != 1 {
t.Fatalf("[DONE] count = %d, want 1", got)
}
if got := len(fake.cancelCallsSnapshot()); got != 1 {
t.Fatalf("initial abort count = %d, want 1", got)
}
}
func TestStreamGateChatTunnelRepeatRecoveryExhaustionTerminatesSSE(t *testing.T) {
safePrefix := uniqueKoreanRunes(520)
service := newScriptedPoolRunService(
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathTunnel), runID: "repeat-tunnel-1",
provider: "provider-a", target: "served-a",
frames: poolChatRepeatSSEFrames("chatcmpl-repeat-1", safePrefix),
},
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathTunnel), runID: "repeat-tunnel-2",
provider: "provider-b", target: "served-b",
frames: poolChatRepeatSSEFrames("chatcmpl-repeat-2", safePrefix),
},
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathTunnel), runID: "repeat-tunnel-3",
provider: "provider-c", target: "served-c",
frames: poolChatRepeatSSEFrames("chatcmpl-repeat-3", safePrefix),
},
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathTunnel), runID: "repeat-tunnel-4",
provider: "provider-d", target: "served-d",
frames: poolChatRepeatSSEFrames("chatcmpl-repeat-4", safePrefix),
},
)
rawBody := []byte(`{"model":"client-model","messages":[{"role":"user","content":"hi"}],"stream":true}`)
srv, dc := buildStreamGatePoolChatFixture(t, service, rawBody, 3)
srv.SetModelCatalog([]config.ModelCatalogEntry{{ID: "client-model", ContextWindowTokens: 32768}})
srv.mu.Lock()
srv.cfg.StreamEvidenceGate.Filters = []config.StreamGateFilterPolicyConf{{
Filter: config.StreamGateFilterRepeatGuard,
Enforcement: config.StreamGateFilterEnforcementBlocking,
Priority: 10,
HoldEvidenceRunes: 500,
}}
srv.mu.Unlock()
w, sink, runErr := runPoolChatStreamGate(t, srv, dc)
if runErr != nil {
t.Fatalf("runtime.Run: %v", runErr)
}
if got := service.poolSubmits() - 1; got != 3 {
t.Fatalf("recovery dispatch count=%d, want 3", got)
}
if got := w.headerCallCount(); got != 1 {
t.Fatalf("WriteHeader count=%d, want 1", got)
}
if w.code != http.StatusOK {
t.Fatalf("status=%d, want %d", w.code, http.StatusOK)
}
composite, ok := sink.(*openAICompositeReleaseSink)
if !ok {
t.Fatalf("release sink=%T, want *openAICompositeReleaseSink", sink)
}
if got := composite.resolvedCodec(); got != openAIStreamGateCodecTunnel {
t.Fatalf("resolved codec=%q, want tunnel", got)
}
if committed, success := composite.terminalStatus(); !committed || success {
t.Fatalf("terminal status=(%v,%v), want (true,false)", committed, success)
}
body := w.body.String()
if got := strings.Count(body, safePrefix); got != 1 {
t.Fatalf("accepted prefix occurrences=%d, want 1", got)
}
if got := strings.Count(body, `"error"`); got != 1 {
t.Fatalf("error envelope count=%d, want 1; body=%q", got, body)
}
if got := strings.Count(body, "data: [DONE]"); got != 1 {
t.Fatalf("[DONE] count=%d, want 1; body=%q", got, body)
}
for _, rejectedID := range []string{
"chatcmpl-repeat-2", "chatcmpl-repeat-3", "chatcmpl-repeat-4",
} {
if strings.Contains(body, rejectedID) {
t.Fatalf("rejected attempt fragment %q reached the caller", rejectedID)
}
}
evidence, err := validateDevRepeatGuardSSE([]byte(body))
if err != nil {
t.Fatalf("caller SSE validation: %v; body=%q", err, body)
}
if evidence.terminal != "error" ||
len(evidence.providerResponseIDs) != 1 ||
evidence.providerResponseIDs[0] != "chatcmpl-repeat-1" {
t.Fatalf("caller SSE evidence=%+v, want initial provider identity and error terminal", evidence)
}
if !strings.HasSuffix(strings.TrimSpace(body), "data: [DONE]") {
t.Fatalf("[DONE] is not the final SSE data event: body=%q", body)
}
}
func TestRepeatGuardStreamOpenNoDuplicatePrefixResponses(t *testing.T) {
safePrefix := uniqueKoreanRunes(520)
initialPayload := responsesStreamPrefix("resp-repeat", "msg-repeat", safePrefix) +
responsesTextDelta("msg-repeat", safePrefix, 5)
service := newScriptedPoolRunService(
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathTunnel), runID: "responses-repeat-initial",
provider: "provider-a", target: "served-a",
frames: responsesTunnelFrames(initialPayload),
},
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathNormalized), runID: "responses-repeat-recovery",
provider: "provider-b", target: "served-b",
runEvents: bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: safePrefix + "recovered suffix"},
&iop.RunEvent{Type: "complete"},
),
},
)
srv, requestCtx, poolReq := buildStreamGateResponsesPoolFixture(t, service, 3)
srv.SetModelCatalog([]config.ModelCatalogEntry{{ID: "client-model", ContextWindowTokens: 8192}})
result, err := service.SubmitProviderPool(context.Background(), poolReq)
if err != nil {
t.Fatalf("initial SubmitProviderPool: %v", err)
}
dc := newOpenAIResponsesPoolTunnelDispatchContext(requestCtx, poolReq)
snapRef, err := requestCtx.ingress.recoveryRef()
if err != nil {
t.Fatalf("recoveryRef: %v", err)
}
gateCfg := config.StreamEvidenceGateConf{
Enabled: true,
Filters: []config.StreamGateFilterPolicyConf{{
Filter: config.StreamGateFilterRepeatGuard,
Enforcement: config.StreamGateFilterEnforcementBlocking,
Priority: 10,
HoldEvidenceRunes: 500,
}},
}
registry, err := openAIStreamGateRegistrySnapshotFor(gateCfg, openAIOutputFilterContext{
endpoint: openAIRebuildEndpointResponses,
requestRef: snapRef.SnapshotRef(),
})
if err != nil {
t.Fatalf("openAIStreamGateRegistrySnapshotFor: %v", err)
}
w := newRecordingResponseWriter()
holder := &openAIResponsesResultHolder{}
selector := newOpenAIStreamGateCodecSelector(openAIStreamGateCodecTunnel)
sink := newOpenAIResponsesPoolReleaseSink(w, holder, selector)
runtime, _, err := srv.buildOpenAIResponsesStreamGateRuntimeFromAttempt(
dc,
openAIAttemptTransport{path: openAIAdmissionTunnel, tunnel: result.Tunnel},
result.Tunnel.Dispatch(),
result.Tunnel.Close,
sink,
registry,
)
if err != nil {
t.Fatalf("buildOpenAIResponsesStreamGateRuntimeFromAttempt: %v", err)
}
if err := runtime.Run(context.Background()); err != nil {
t.Fatalf("runtime.Run: %v", err)
}
if err := runtime.CloseRequestResources(context.Background(), true); err != nil {
t.Fatalf("CloseRequestResources: %v", err)
}
if got := service.poolSubmits(); got != 2 {
t.Fatalf("SubmitProviderPool calls = %d, want initial plus one recovery", got)
}
payloads := assertResponsesSSELifecycle(t, w.body.String(), "response.completed")
if got := responsesSSEOutputText(payloads); got != safePrefix+"recovered suffix" {
t.Fatalf("Responses downstream content length=%d, want safe prefix once plus recovered suffix", len(got))
}
responseStarts := 0
terminals := 0
for _, payload := range payloads {
switch payload["type"] {
case "response.created":
responseStarts++
case "response.completed", "error":
terminals++
}
}
if responseStarts != 1 || terminals != 1 {
t.Fatalf("Responses lifecycle starts=%d terminals=%d, want 1/1", responseStarts, terminals)
}
}
func TestOpenAIRepeatResumeContextOverflowDoesNotDispatch(t *testing.T) {
srv, dc, initial, fake := buildStreamGateChatFixture(t, true, 3)
// No catalog entry means the context window is unknown. S20 must refuse the
// continuation before the dispatcher can submit a replacement attempt.
initial.events = bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "output awaiting a repeat resume"},
&iop.RunEvent{Type: "complete"},
)
snapRef, err := dc.ingress.recoveryRef()
if err != nil {
t.Fatalf("recoveryRef: %v", err)
}
violation := newInjectedContinuationViolationFilter(t, "test.resume.unknown-context", 1, snapRef.SnapshotRef(), len("output "))
reg, err := streamgate.NewFilterRegistration(violation, streamGateNoopCapability, true, streamgate.FilterEnforcementBlocking, streamGateFilterTimeout, injectedViolationPriority)
if err != nil {
t.Fatalf("NewFilterRegistration: %v", err)
}
w := newRecordingResponseWriter()
rt, _, err := srv.buildOpenAIChatStreamGateRuntime(dc, initial, newOpenAIChatSSEReleaseSink(w, w, "chatcmpl-run-attempt-1", time.Now().Unix(), "client-model"), streamGateTestRegistry(t, reg))
if err != nil {
t.Fatalf("buildOpenAIChatStreamGateRuntime: %v", err)
}
if err := rt.Run(context.Background()); err != nil {
t.Fatalf("rt.Run: %v", err)
}
if got := len(fake.reqsSnapshot()); got != 0 {
t.Fatalf("context refusal dispatched %d recovery attempts, want 0", got)
}
}
func buildStreamGateResponsesFixture(t *testing.T, service runService, providerPool bool) (*Server, *responsesDispatchContext) {
t.Helper()
body := []byte(`{"model":"client-model","input":"CALLER_RESPONSES_INPUT_MUST_NOT_BE_COPIED"}`)
fault := 3
srv := NewServer(config.EdgeOpenAIConf{
Adapter: "ollama", Target: "llama-fixed", SessionID: "cli", TimeoutSec: 15,
StreamEvidenceGate: config.StreamEvidenceGateConf{Enabled: true, MaxRequestFaultRecovery: &fault},
}, service, nil)
srv.SetModelCatalog([]config.ModelCatalogEntry{{ID: "client-model", ContextWindowTokens: 4096}})
route := routeDispatch{Adapter: "ollama", Target: "llama-fixed", SessionID: "cli", TimeoutSec: 15, ProviderPool: providerPool}
base := newTestRequestContext(t, route, body)
base.r = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
base.endpoint = usageEndpointResponses
requestCtx := &responsesRequestContext{openAIRequestContext: base, envelope: responsesEnvelope{Model: "client-model"}}
dc, err := srv.newResponsesDispatchContext(requestCtx, responsesRequest{Model: "client-model", Input: json.RawMessage(`"CALLER_RESPONSES_INPUT_MUST_NOT_BE_COPIED"`)})
if err != nil {
t.Fatalf("newResponsesDispatchContext: %v", err)
}
if !providerPool {
return srv, dc
}
pool := edgeservice.ProviderPoolDispatchRequest{
Run: dc.submitReq,
Tunnel: edgeservice.SubmitProviderTunnelRequest{
ModelGroupKey: "client-model", SessionID: route.SessionID, Method: http.MethodPost,
Path: "/v1/responses", ProviderPool: true,
},
}
return srv, dc.withPoolDispatch(pool)
}
func TestOpenAIResponsesRepeatResumeDispatchesNormalized(t *testing.T) {
fake := &fakeRunService{eventRuns: []chan *iop.RunEvent{bufferedRunEvents(&iop.RunEvent{Type: "complete"})}, runIDs: []string{"responses-attempt-2"}}
srv, dc := buildStreamGateResponsesFixture(t, fake, false)
source := newOpenAIRecoverySourceStore(dc.ingress)
source.recordText("safe prefix repeated tail")
ref, err := dc.ingress.recoveryRef()
if err != nil {
t.Fatalf("recoveryRef: %v", err)
}
rebuilder, err := newOpenAIRequestRebuilder(dc.ingress, openAIRebuildEndpointResponses, source, 4096)
if err != nil {
t.Fatalf("newOpenAIRequestRebuilder: %v", err)
}
defer rebuilder.Close()
directive, err := streamgate.NewRecoveryDirectiveContinuation(len("safe prefix "), ref.SnapshotRef())
if err != nil {
t.Fatalf("NewRecoveryDirectiveContinuation: %v", err)
}
plan := mustOpenAIRecoveryPlan(t, "plan.responses.dispatch.normalized", streamgate.RecoveryStrategyContinuationRepair, directive)
draft, err := rebuilder.RebuildRequest(context.Background(), ref, plan)
if err != nil {
t.Fatalf("RebuildRequest: %v", err)
}
_, request, err := plan.FinalizeRebuiltRequest(draft)
if err != nil {
t.Fatalf("FinalizeRebuiltRequest: %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, fmt.Errorf("normalized Responses replacement was not admitted")
}
return newOpenAIResponsesEventSource(state.get(), transport.run, &openAIResponsesResultHolder{}, nil), nil
})
if err != nil {
t.Fatalf("newOpenAIAttemptDispatcher: %v", err)
}
binding, err := dispatcher.DispatchAttempt(context.Background(), request)
if err != nil {
t.Fatalf("DispatchAttempt: %v", err)
}
for {
event, eventErr := binding.EventSource().NextEvent(context.Background())
if eventErr != nil {
t.Fatalf("replacement event source: %v", eventErr)
}
if event.Kind() == streamgate.EventKindTerminal {
break
}
}
requests := fake.reqsSnapshot()
if len(requests) != 1 {
t.Fatalf("recovery SubmitRun calls = %d, want 1", len(requests))
}
if strings.Contains(requests[0].Prompt, "CALLER_RESPONSES_INPUT_MUST_NOT_BE_COPIED") {
t.Fatalf("recovery request copied caller input: %q", requests[0].Prompt)
}
if got := requests[0].Input["responses_resume_content"]; got != "safe prefix " {
t.Fatalf("recovery content provenance = %#v", got)
}
}
func TestOpenAIResponsesRepeatResumeDispatchesProviderPool(t *testing.T) {
service := newScriptedPoolRunService(
scriptedPoolAttempt{path: string(edgeservice.ProviderPoolPathNormalized), runID: "responses-pool-attempt-2", provider: "provider-b", target: "served-b", runEvents: bufferedRunEvents(
&iop.RunEvent{Type: "complete"},
)},
)
srv, dc := buildStreamGateResponsesFixture(t, service, true)
source := newOpenAIRecoverySourceStore(dc.ingress)
source.recordText("safe prefix repeated tail")
ref, err := dc.ingress.recoveryRef()
if err != nil {
t.Fatalf("recoveryRef: %v", err)
}
rebuilder, err := newOpenAIRequestRebuilder(dc.ingress, openAIRebuildEndpointResponses, source, 4096)
if err != nil {
t.Fatalf("newOpenAIRequestRebuilder: %v", err)
}
defer rebuilder.Close()
directive, err := streamgate.NewRecoveryDirectiveContinuation(len("safe prefix "), ref.SnapshotRef())
if err != nil {
t.Fatalf("NewRecoveryDirectiveContinuation: %v", err)
}
plan := mustOpenAIRecoveryPlan(t, "plan.responses.dispatch.pool", streamgate.RecoveryStrategyContinuationRepair, directive)
draft, err := rebuilder.RebuildRequest(context.Background(), ref, plan)
if err != nil {
t.Fatalf("RebuildRequest: %v", err)
}
_, request, err := plan.FinalizeRebuiltRequest(draft)
if err != nil {
t.Fatalf("FinalizeRebuiltRequest: %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, fmt.Errorf("provider-pool Responses replacement was not admitted as normalized")
}
return newOpenAIResponsesEventSource(state.get(), transport.run, &openAIResponsesResultHolder{}, nil), nil
})
if err != nil {
t.Fatalf("newOpenAIAttemptDispatcher: %v", err)
}
binding, err := dispatcher.DispatchAttempt(context.Background(), request)
if err != nil {
t.Fatalf("DispatchAttempt: %v", err)
}
for {
event, eventErr := binding.EventSource().NextEvent(context.Background())
if eventErr != nil {
t.Fatalf("provider-pool replacement event source: %v", eventErr)
}
if event.Kind() == streamgate.EventKindTerminal {
break
}
}
pools, _, _, _, runRequests, _ := service.snapshot()
if pools != 1 || len(runRequests) != 1 {
t.Fatalf("provider-pool replacement = pools:%d runs:%d, want 1 admission", pools, len(runRequests))
}
resume := runRequests[0]
if !resume.ProviderPool || strings.Contains(resume.Prompt, "CALLER_RESPONSES_INPUT_MUST_NOT_BE_COPIED") {
t.Fatalf("pool resume request = %#v", resume)
}
if got := resume.Input["responses_resume_content"]; got != "safe prefix " {
t.Fatalf("pool recovery content provenance = %#v", got)
}
}
func TestStreamGateChatBuildRuntimeRecoveryBudgetExhaustionTerminatesWithError(t *testing.T) {
srv, dc, initial, fake := buildStreamGateChatFixture(t, true, 1)
fake.eventRuns = []chan *iop.RunEvent{bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "still-bad"},
&iop.RunEvent{Type: "complete"},
)}
fake.runIDs = []string{"run-attempt-2"}
snapRef, err := dc.ingress.recoveryRef()
if err != nil {
t.Fatalf("recoveryRef: %v", err)
}
// trigger=2: violates on both the initial attempt and the one allowed
// recovery attempt, so the request-total cap (1) is exhausted after the
// first recovery and the second violation must terminate, not dispatch a
// third attempt.
violation := newInjectedViolationFilter(t, "test.injected_violation", 2, snapRef.SnapshotRef())
reg, err := streamgate.NewFilterRegistration(violation, streamGateNoopCapability, true, streamgate.FilterEnforcementBlocking, streamGateFilterTimeout, injectedViolationPriority)
if err != nil {
t.Fatalf("NewFilterRegistration: %v", err)
}
registry := streamGateTestRegistry(t, reg)
w := newRecordingResponseWriter()
rt, _, err := srv.buildOpenAIChatStreamGateRuntime(dc, initial, newOpenAIChatSSEReleaseSink(w, w, "chatcmpl-run-attempt-1", time.Now().Unix(), "client-model"), registry)
if err != nil {
t.Fatalf("buildOpenAIChatStreamGateRuntime: %v", err)
}
if err := rt.Run(context.Background()); err != nil {
t.Fatalf("rt.Run: %v", err)
}
if got := len(fake.reqsSnapshot()); got != 1 {
t.Fatalf("recovery dispatch count: got %d, want 1 (cap=1 allows exactly one recovery dispatch)", got)
}
if !strings.Contains(w.body.String(), `"error"`) {
t.Fatalf("expected an SSE error terminal, got body=%s", w.body.String())
}
if strings.Contains(w.body.String(), "still-bad") {
t.Fatalf("exhausted recovery must not leak the still-violating attempt's content: body=%s", w.body.String())
}
}
func TestStreamGateChatBuildRuntimeObserveOnlyDoesNotBlock(t *testing.T) {
srv, dc, initial, fake := buildStreamGateChatFixture(t, true, 3)
initial.events = bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "observed content"},
&iop.RunEvent{Type: "complete"},
)
snapRef, err := dc.ingress.recoveryRef()
if err != nil {
t.Fatalf("recoveryRef: %v", err)
}
violation := newInjectedViolationFilter(t, "test.injected_violation_observe", 1, snapRef.SnapshotRef())
reg, err := streamgate.NewFilterRegistration(violation, streamGateNoopCapability, true, streamgate.FilterEnforcementObserveOnly, streamGateFilterTimeout, injectedViolationPriority)
if err != nil {
t.Fatalf("NewFilterRegistration: %v", err)
}
registry := streamGateTestRegistry(t, reg)
w := newRecordingResponseWriter()
rt, _, err := srv.buildOpenAIChatStreamGateRuntime(dc, initial, newOpenAIChatSSEReleaseSink(w, w, "chatcmpl-run-attempt-1", time.Now().Unix(), "client-model"), registry)
if err != nil {
t.Fatalf("buildOpenAIChatStreamGateRuntime: %v", err)
}
if err := rt.Run(context.Background()); err != nil {
t.Fatalf("rt.Run: %v", err)
}
if got := len(fake.reqsSnapshot()); got != 0 {
t.Fatalf("observe-only violation must not trigger a recovery dispatch: got %d dispatches", got)
}
chunks := parseSSEChatChunks(t, w.body.String())
if got := joinedContent(chunks); got != "observed content" {
t.Fatalf("joined content: got %q, want the pass-through content unblocked by the observe-only violation", got)
}
}
// --- adapter-level: tunnel runtime Noop passthrough ---------------------------
func TestStreamGateTunnelBuildRuntimeNoopPassRelaysRawBytes(t *testing.T) {
rawBody := []byte(`{"model":"client-model","messages":[{"role":"user","content":"hi"}],"stream":true}`)
fault := 3
fake := &providerFakeRunService{}
srv := NewServer(config.EdgeOpenAIConf{
Adapter: "openai-compat",
Target: "served-model",
TimeoutSec: 15,
StreamEvidenceGate: config.StreamEvidenceGateConf{
Enabled: true,
MaxRequestFaultRecovery: &fault,
},
}, fake, nil)
route := routeDispatch{Adapter: "openai-compat", Target: "served-model", TimeoutSec: 15}
requestCtx := newTestRequestContext(t, route, rawBody)
frames := make(chan *iop.ProviderTunnelFrame, 4)
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_RESPONSE_START, StatusCode: 200, Headers: map[string]string{"Content-Type": "text/event-stream"}}
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: []byte("data: {\"choices\":[{\"delta\":{\"content\":\"raw\"}}]}\n\n")}
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_END}
close(frames)
handle := &fakeTunnelHandle{
dispatch: edgeservice.RunDispatch{RunID: "tunnel-attempt-1", Adapter: "openai-compat", Target: "served-model"},
frames: frames,
}
req := openAITunnelStreamGateRequest{
route: route,
ingress: requestCtx.ingress,
endpoint: openAIRebuildEndpointChat,
method: http.MethodPost,
path: "/v1/chat/completions",
modelGroupKey: "client-model",
requestModel: "", // no model echo present in the fixture body; rewrite is a no-op
authorize: func(context.Context) (map[string]string, error) { return nil, nil },
rewriteBody: func(body []byte, target string) ([]byte, error) {
return body, nil
},
}
w := newRecordingResponseWriter()
registry := streamGateTestRegistry(t)
rt, _, err := srv.buildOpenAITunnelStreamGateRuntime(req, handle, newOpenAITunnelReleaseSink(w, w), registry)
if err != nil {
t.Fatalf("buildOpenAITunnelStreamGateRuntime: %v", err)
}
if err := rt.Run(context.Background()); err != nil {
t.Fatalf("rt.Run: %v", err)
}
if w.code != http.StatusOK {
t.Fatalf("status: got %d", w.code)
}
if got := w.headerCallCount(); got != 1 {
t.Fatalf("WriteHeader call count: got %d, want 1", got)
}
if !strings.Contains(w.body.String(), `"content":"raw"`) {
t.Fatalf("raw provider bytes were not relayed byte-for-byte: body=%s", w.body.String())
}
}
// --- S16: strategy recovery cap enforcement (REVIEW_EDGE_VERTICAL_SLICE-4) ---
// newStrategyCapChatFixture mirrors buildStreamGateChatFixture but pins an
// explicit per-strategy recovery cap so the request-start RuntimeOptions
// snapshot carries a non-empty strategy-limit map into the Core policy.
func newStrategyCapChatFixture(t *testing.T, maxRequestFaultRecovery, maxStrategyFaultRecovery int) (*Server, *chatDispatchContext, *fakeRunResult, *fakeRunService) {
t.Helper()
rawBody := []byte(`{"model":"client-model","messages":[{"role":"user","content":"hi"}],"stream":true}`)
var req chatCompletionRequest
if err := json.Unmarshal(rawBody, &req); err != nil {
t.Fatalf("decode fixture body: %v", err)
}
total := maxRequestFaultRecovery
strategy := maxStrategyFaultRecovery
fake := &fakeRunService{}
srv := NewServer(config.EdgeOpenAIConf{
Adapter: "ollama",
Target: "llama-fixed",
SessionID: "cli",
TimeoutSec: 15,
StreamEvidenceGate: config.StreamEvidenceGateConf{
Enabled: true,
MaxRequestFaultRecovery: &total,
MaxStrategyFaultRecovery: &strategy,
},
}, fake, nil)
route := routeDispatch{Adapter: "ollama", Target: "llama-fixed", SessionID: "cli", TimeoutSec: 15}
requestCtx := newTestRequestContext(t, route, rawBody)
basePrompt := promptFromMessages(req.Messages)
outputPolicy := srv.resolveOutputPolicy(basePrompt)
dc := srv.newChatDispatchContext(requestCtx, req, basePrompt, outputPolicy)
initial := &fakeRunResult{
dispatch: edgeservice.RunDispatch{RunID: "run-attempt-1", Adapter: "ollama", Target: "llama-fixed", SessionID: "cli", TimeoutSec: 15},
events: bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "should-be-discarded"},
&iop.RunEvent{Type: "complete"},
),
}
return srv, dc, initial, fake
}
func runStrategyCapChatFixture(t *testing.T, srv *Server, dc *chatDispatchContext, initial *fakeRunResult, trigger int) (*recordingResponseWriter, error) {
t.Helper()
snapRef, err := dc.ingress.recoveryRef()
if err != nil {
t.Fatalf("recoveryRef: %v", err)
}
violation := newInjectedViolationFilter(t, "test.injected_violation", trigger, snapRef.SnapshotRef())
reg, err := streamgate.NewFilterRegistration(violation, streamGateNoopCapability, true, streamgate.FilterEnforcementBlocking, streamGateFilterTimeout, injectedViolationPriority)
if err != nil {
t.Fatalf("NewFilterRegistration: %v", err)
}
registry := streamGateTestRegistry(t, reg)
w := newRecordingResponseWriter()
rt, _, err := srv.buildOpenAIChatStreamGateRuntime(dc, initial, newOpenAIChatSSEReleaseSink(w, w, "chatcmpl-run-attempt-1", time.Now().Unix(), "client-model"), registry)
if err != nil {
t.Fatalf("buildOpenAIChatStreamGateRuntime: %v", err)
}
return w, rt.Run(context.Background())
}
// TestStreamGateChatStrategyCapZeroBlocksExactReplayRecovery verifies that an
// explicit strategy cap of 0 forbids the exact_replay fault strategy even when
// the request-total cap allows recovery: the single violation must terminate
// with an error and dispatch zero recovery attempts.
func TestStreamGateChatStrategyCapZeroBlocksExactReplayRecovery(t *testing.T) {
srv, dc, initial, fake := newStrategyCapChatFixture(t, 3, 0)
fake.eventRuns = []chan *iop.RunEvent{bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "must-not-be-reached"},
&iop.RunEvent{Type: "complete"},
)}
fake.runIDs = []string{"run-attempt-2"}
w, runErr := runStrategyCapChatFixture(t, srv, dc, initial, 1)
if runErr != nil {
t.Fatalf("rt.Run: %v", runErr)
}
if got := len(fake.reqsSnapshot()); got != 0 {
t.Fatalf("strategy cap 0 must block all exact_replay dispatch: got %d dispatches", got)
}
if !strings.Contains(w.body.String(), `"error"`) {
t.Fatalf("expected an SSE error terminal for a strategy-capped violation, body=%s", w.body.String())
}
}
// TestStreamGateChatStrategyCapOneExhaustsOnSecondViolation verifies that a
// strategy cap of 1 (below the request-total cap of 3) allows exactly one
// exact_replay recovery dispatch and then exhausts, terminating on the second
// violation instead of dispatching a second recovery.
func TestStreamGateChatStrategyCapOneExhaustsOnSecondViolation(t *testing.T) {
srv, dc, initial, fake := newStrategyCapChatFixture(t, 3, 1)
fake.eventRuns = []chan *iop.RunEvent{bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "still-bad"},
&iop.RunEvent{Type: "complete"},
)}
fake.runIDs = []string{"run-attempt-2"}
w, runErr := runStrategyCapChatFixture(t, srv, dc, initial, 2)
if runErr != nil {
t.Fatalf("rt.Run: %v", runErr)
}
if got := len(fake.reqsSnapshot()); got != 1 {
t.Fatalf("strategy cap 1 allows exactly one exact_replay dispatch: got %d", got)
}
if !strings.Contains(w.body.String(), `"error"`) {
t.Fatalf("expected an SSE error terminal once the strategy cap is exhausted, body=%s", w.body.String())
}
}
// --- S16: current-binding lifecycle + terminal outcome (REVIEW_EDGE_VERTICAL_SLICE-3) ---
// TestStreamGateChatCloseRequestResourcesAbortsLatestRun verifies that after a
// recovery cycle the request runtime cleans up the *current* (latest) attempt
// binding — not just the initial handle — and that a non-graceful close cancels
// the latest provider run while a graceful close does not, and that repeated
// closes are idempotent.
func TestStreamGateChatCloseRequestResourcesAbortsLatestRun(t *testing.T) {
for _, graceful := range []bool{false, true} {
graceful := graceful
name := "abort"
if graceful {
name = "graceful"
}
t.Run(name, func(t *testing.T) {
srv, dc, initial, fake := buildStreamGateChatFixture(t, true, 3)
fake.eventRuns = []chan *iop.RunEvent{bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "final answer"},
&iop.RunEvent{Type: "complete"},
)}
fake.runIDs = []string{"run-attempt-2"}
snapRef, err := dc.ingress.recoveryRef()
if err != nil {
t.Fatalf("recoveryRef: %v", err)
}
violation := newInjectedViolationFilter(t, "test.injected_violation", 1, snapRef.SnapshotRef())
reg, err := streamgate.NewFilterRegistration(violation, streamGateNoopCapability, true, streamgate.FilterEnforcementBlocking, streamGateFilterTimeout, injectedViolationPriority)
if err != nil {
t.Fatalf("NewFilterRegistration: %v", err)
}
registry := streamGateTestRegistry(t, reg)
w := newRecordingResponseWriter()
rt, _, err := srv.buildOpenAIChatStreamGateRuntime(dc, initial, newOpenAIChatSSEReleaseSink(w, w, "chatcmpl-run-attempt-1", time.Now().Unix(), "client-model"), registry)
if err != nil {
t.Fatalf("buildOpenAIChatStreamGateRuntime: %v", err)
}
if err := rt.Run(context.Background()); err != nil {
t.Fatalf("rt.Run: %v", err)
}
// The recovery cycle already aborted the initial attempt (run-attempt-1).
cancelsAfterRun := fake.cancelCallsSnapshot()
if !cancelRunIDsContain(cancelsAfterRun, "run-attempt-1") {
t.Fatalf("recovery must abort the initial attempt run-attempt-1, cancels=%v", cancelRunIDs(cancelsAfterRun))
}
if cancelRunIDsContain(cancelsAfterRun, "run-attempt-2") {
t.Fatalf("the latest attempt must not be canceled by a successful terminal, cancels=%v", cancelRunIDs(cancelsAfterRun))
}
if err := rt.CloseRequestResources(context.Background(), graceful); err != nil {
t.Fatalf("CloseRequestResources: %v", err)
}
// Idempotent: a second close must be a no-op regardless of the flag.
if err := rt.CloseRequestResources(context.Background(), !graceful); err != nil {
t.Fatalf("second CloseRequestResources: %v", err)
}
cancelsAfterClose := fake.cancelCallsSnapshot()
latestCanceled := cancelRunIDsContain(cancelsAfterClose, "run-attempt-2")
if graceful && latestCanceled {
t.Fatalf("graceful close must not cancel the latest run, cancels=%v", cancelRunIDs(cancelsAfterClose))
}
if !graceful && !latestCanceled {
t.Fatalf("non-graceful close must cancel the latest run once, cancels=%v", cancelRunIDs(cancelsAfterClose))
}
// Exactly-once: at most one cancel per attempt id.
if n := cancelRunIDCount(cancelsAfterClose, "run-attempt-1"); n != 1 {
t.Fatalf("initial attempt canceled %d times, want exactly 1", n)
}
if n := cancelRunIDCount(cancelsAfterClose, "run-attempt-2"); n > 1 {
t.Fatalf("latest attempt canceled %d times, want at most 1", n)
}
})
}
}
func cancelRunIDs(cancels []edgeservice.CancelRunRequest) []string {
ids := make([]string, 0, len(cancels))
for _, c := range cancels {
ids = append(ids, c.RunID)
}
return ids
}
func cancelRunIDsContain(cancels []edgeservice.CancelRunRequest, runID string) bool {
return cancelRunIDCount(cancels, runID) > 0
}
func cancelRunIDCount(cancels []edgeservice.CancelRunRequest, runID string) int {
n := 0
for _, c := range cancels {
if c.RunID == runID {
n++
}
}
return n
}
// TestStreamGateUsageStatusFromTerminalOutcome verifies the host maps a
// Core-committed terminal into the correct usage-metric status: a committed
// error terminal is reported as error (never success), a committed success
// terminal as success, and a caller cancel/timeout as cancel.
func TestStreamGateUsageStatusFromTerminalOutcome(t *testing.T) {
cases := []struct {
name string
runErr error
committed bool
success bool
want string
}{
{"success terminal", nil, true, true, usageStatusSuccess},
{"error terminal not success metric", nil, true, false, usageStatusError},
{"caller cancel", context.Canceled, false, false, usageStatusCancel},
{"deadline exceeded", context.DeadlineExceeded, false, false, usageStatusCancel},
{"internal run error", fmt.Errorf("boom"), false, false, usageStatusError},
{"no terminal committed", nil, false, false, usageStatusError},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
if got := streamGateUsageStatus(tc.runErr, tc.committed, tc.success); got != tc.want {
t.Fatalf("streamGateUsageStatus(%v,%v,%v) = %q, want %q", tc.runErr, tc.committed, tc.success, got, tc.want)
}
})
}
}
// TestOpenAIChatSSESinkCapturesTerminalOutcome verifies the release sink records
// the terminal disposition the Core commits so the host outcome is authoritative.
func TestOpenAIChatSSESinkCapturesTerminalOutcome(t *testing.T) {
successTerm, err := streamgate.NewSuccessTerminalResult(streamGateChannelDefault, time.Now())
if err != nil {
t.Fatalf("NewSuccessTerminalResult: %v", err)
}
w := newRecordingResponseWriter()
sink := newOpenAIChatSSEReleaseSink(w, w, "chatcmpl-x", time.Now().Unix(), "m")
if committed, _ := sink.terminalStatus(); committed {
t.Fatalf("sink reported a committed terminal before any commit")
}
if _, err := sink.CommitTerminal(context.Background(), successTerm); err != nil {
t.Fatalf("CommitTerminal: %v", err)
}
committed, success := sink.terminalStatus()
if !committed || !success {
t.Fatalf("sink terminalStatus after success = (%v,%v), want (true,true)", committed, success)
}
}
// Note: the separate fixed-run recovery and fixed-tunnel pass fixtures above
// do not prove the S16 provider/path-switch requirement. That evidence needs a
// provider-pool host fixture in which one RequestRuntime recovery cycle changes
// the actual AttemptBinding execution path before release.
// --- S16: provider-pool one cycle, path switch, and single re-admission ------
// buildStreamGatePoolChatFixture builds a provider-pool chat dispatch context
// bound to a scripted pool admission double, mirroring what
// handleChatCompletionsProviderPool freezes before its first SubmitProviderPool
// call. The returned pool request is the same template the recovery admission
// builder re-enters for every recovery attempt.
func buildStreamGatePoolChatFixture(t *testing.T, service runService, rawBody []byte, maxRequestFaultRecovery int) (*Server, *chatDispatchContext) {
t.Helper()
var req chatCompletionRequest
if err := json.Unmarshal(rawBody, &req); err != nil {
t.Fatalf("decode fixture body: %v", err)
}
fault := maxRequestFaultRecovery
srv := NewServer(config.EdgeOpenAIConf{
Adapter: "ollama",
Target: "llama-fixed",
TimeoutSec: 15,
StreamEvidenceGate: config.StreamEvidenceGateConf{
Enabled: true,
MaxRequestFaultRecovery: &fault,
},
}, service, nil)
route := routeDispatch{ProviderPool: true, SessionID: "cli", TimeoutSec: 15}
requestCtx := newTestRequestContext(t, route, rawBody)
basePrompt := promptFromMessages(req.Messages)
outputPolicy := srv.resolveOutputPolicy(basePrompt)
dc := srv.newChatDispatchContext(requestCtx, req, basePrompt, outputPolicy)
poolReq := edgeservice.ProviderPoolDispatchRequest{
Run: dc.submitReq,
Tunnel: edgeservice.SubmitProviderTunnelRequest{
ModelGroupKey: strings.TrimSpace(req.Model),
SessionID: route.SessionID,
Method: http.MethodPost,
Path: "/v1/chat/completions",
Stream: req.Stream,
TimeoutSec: route.TimeoutSec,
Metadata: dc.runMetadata,
EstimatedInputTokens: dc.estimate,
ContextClass: dc.contextClass,
ProviderPool: true,
},
PrepareTunnel: func(tunnelReq edgeservice.SubmitProviderTunnelRequest) (edgeservice.SubmitProviderTunnelRequest, error) {
return tunnelReq, nil
},
PrepareRun: func(runReq edgeservice.SubmitRunRequest) (edgeservice.SubmitRunRequest, error) {
return runReq, nil
},
}
poolReq.Tunnel.BuildBody = func(target string) ([]byte, error) {
body, err := dc.ingress.canonicalBody()
if err != nil {
return nil, err
}
return rewriteChatCompletionModel(body, target, req)
}
return srv, dc.withPoolDispatch(poolReq)
}
// runPoolChatStreamGate dispatches the initial pool admission exactly the way
// handleChatCompletionsProviderPool does, then drives the production
// provider-pool wiring (config, composite sink, dual event-source factory,
// recovery admission builder) through one request runtime.
func runPoolChatStreamGate(t *testing.T, srv *Server, dc *chatDispatchContext, extra ...streamgate.FilterRegistration) (*recordingResponseWriter, openAIStreamGateSink, error) {
t.Helper()
result, err := srv.service.SubmitProviderPool(context.Background(), *dc.poolDispatch)
if err != nil {
t.Fatalf("initial SubmitProviderPool: %v", err)
}
w := newRecordingResponseWriter()
cfg, sink, err := srv.newOpenAIChatPoolStreamGateConfig(w, w, dc, result, extra)
if err != nil {
t.Fatalf("newOpenAIChatPoolStreamGateConfig: %v", err)
}
rt, _, err := srv.buildOpenAIChatStreamGateRuntimeFor(dc, cfg)
if err != nil {
t.Fatalf("buildOpenAIChatStreamGateRuntimeFor: %v", err)
}
runErr := rt.Run(context.Background())
_ = rt.CloseRequestResources(context.Background(), false)
return w, sink, runErr
}
func newInjectedViolationRegistration(t *testing.T, dc *chatDispatchContext, id string, trigger int) streamgate.FilterRegistration {
t.Helper()
snapRef, err := dc.ingress.recoveryRef()
if err != nil {
t.Fatalf("recoveryRef: %v", err)
}
filter := newInjectedViolationFilter(t, id, trigger, snapRef.SnapshotRef())
reg, err := streamgate.NewFilterRegistration(filter, streamGateNoopCapability, true, streamgate.FilterEnforcementBlocking, streamGateFilterTimeout, injectedViolationPriority)
if err != nil {
t.Fatalf("NewFilterRegistration: %v", err)
}
return reg
}
func poolChatSSEFrames(content string) chan *iop.ProviderTunnelFrame {
return bufferedTunnelFrames(
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_RESPONSE_START, StatusCode: 200, Headers: map[string]string{"Content-Type": "text/event-stream"}},
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: []byte("data: {\"choices\":[{\"delta\":{\"content\":\"" + content + "\"}}]}\n\n")},
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_END, End: true},
)
}
func poolChatRepeatSSEFrames(responseID, content string) chan *iop.ProviderTunnelFrame {
event := func(value string) []byte {
return []byte(fmt.Sprintf(
"data: {\"id\":%q,\"object\":\"chat.completion.chunk\",\"choices\":[{\"index\":0,\"delta\":{\"content\":%q},\"finish_reason\":null}]}\n\n",
responseID, value,
))
}
return bufferedTunnelFrames(
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_RESPONSE_START, StatusCode: 200, Headers: map[string]string{"Content-Type": "text/event-stream"}},
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: event(content)},
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: event(content)},
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_END, End: true},
)
}
type recordingOpenAIObservationSink struct {
mu sync.Mutex
got []streamgate.FilterObservation
delegate streamgate.ObservationSink
fail error
}
func (s *recordingOpenAIObservationSink) Emit(ctx context.Context, obs streamgate.FilterObservation) error {
s.mu.Lock()
s.got = append(s.got, obs)
s.mu.Unlock()
if s.delegate != nil {
if err := s.delegate.Emit(ctx, obs); err != nil {
return err
}
}
return s.fail
}
func (s *recordingOpenAIObservationSink) Snapshot() []streamgate.FilterObservation {
s.mu.Lock()
defer s.mu.Unlock()
return append([]streamgate.FilterObservation(nil), s.got...)
}
func observationKinds(observations []streamgate.FilterObservation) []streamgate.ObservationKind {
kinds := make([]streamgate.ObservationKind, 0, len(observations))
for _, obs := range observations {
kinds = append(kinds, obs.Kind())
}
return kinds
}
func assertProviderPoolObservationTimeline(t *testing.T, observations []streamgate.FilterObservation) {
t.Helper()
if len(observations) == 0 {
t.Fatal("observation timeline is empty")
}
initialTarget, err := streamgate.NewObservationAttemptTarget("client-model", "served-normalized", "prov-normalized", "normalized")
if err != nil {
t.Fatalf("initial target: %v", err)
}
reboundTarget, err := streamgate.NewObservationAttemptTarget("client-model", "served-tunnel", "prov-tunnel", "provider_tunnel")
if err != nil {
t.Fatalf("rebound target: %v", err)
}
// Per-observation invariants over the whole timeline: every row is raw-free
// valid, sequence is strictly monotonic from 1, correlation is stable, causes
// and evidence stay bounded, and the actual attempt id/target matches the
// segment the row belongs to. The segment switches to the rebound attempt at
// recovery_dispatched (inclusive), so any extra interleaved filter row is
// still attributed to the correct actual attempt.
seenDispatch := false
for i, obs := range observations {
if err := obs.Validate(); err != nil {
t.Fatalf("observation[%d] invalid: %v", i, err)
}
if obs.Sequence() != uint64(i+1) {
t.Fatalf("observation[%d] sequence=%d, want %d", i, obs.Sequence(), i+1)
}
if obs.StableCorrelation() != "req.pool-attempt-1" || obs.ConfigGeneration() != streamGateConfigGeneration {
t.Fatalf("observation[%d] correlation=(%q,%q)", i, obs.StableCorrelation(), obs.ConfigGeneration())
}
if obs.Kind() == streamgate.ObservationKindRecoveryDispatched {
seenDispatch = true
}
wantID, wantTarget := "attempt.pool-attempt-1", initialTarget
if seenDispatch {
wantID, wantTarget = "attempt.pool-attempt-2", reboundTarget
}
if obs.AttemptID() != wantID || obs.AttemptTarget() != wantTarget {
t.Fatalf("observation[%d] kind=%s attempt/target=(%q,%+v), want (%q,%+v)", i, obs.Kind(), obs.AttemptID(), obs.AttemptTarget(), wantID, wantTarget)
}
if obs.Causes().Len() > streamgate.MaxFailureCauses {
t.Fatalf("observation[%d] cause count=%d", i, obs.Causes().Len())
}
if evidence := obs.Evidence(); evidence != nil {
if err := evidence.Validate(); err != nil {
t.Fatalf("observation[%d] evidence invalid: %v", i, err)
}
}
}
// Exact chronological order: the actual normalized→tunnel recovery must
// appear as this ordered subsequence, consumed front-to-back. Extra filter
// rows (one evaluation_started/evaluated pair per registered filter) may
// interleave, but the relative order of these required rows must hold — a
// swapped recovery/terminal or a rebound segment that starts before dispatch
// can no longer pass. Each row pins the actual attempt segment, the commit
// state, and per-segment epoch consistency.
const (
segInitial = "initial"
segRebound = "rebound"
)
type wantRow struct {
kind streamgate.ObservationKind
segment string
commit streamgate.CommitState
// epochGroup buckets rows that must share the same non-zero epoch. The
// recovery rows carry the triggering (initial evaluation) epoch, so they
// share the initial group. An empty group means the row is epoch-unbound.
epochGroup string
}
uncommitted := streamgate.CommitStateTransportUncommitted
rows := []wantRow{
{streamgate.ObservationKindResponseStaged, segInitial, uncommitted, ""},
{streamgate.ObservationKindFilterEvaluationStarted, segInitial, uncommitted, "eval-initial"},
{streamgate.ObservationKindFilterEvaluated, segInitial, uncommitted, "eval-initial"},
{streamgate.ObservationKindArbitrationDecided, segInitial, uncommitted, "eval-initial"},
{streamgate.ObservationKindRecoveryPlanSelected, segInitial, uncommitted, "eval-initial"},
{streamgate.ObservationKindRecoveryAttemptAborted, segInitial, uncommitted, "eval-initial"},
{streamgate.ObservationKindRecoveryRebuilt, segInitial, uncommitted, "eval-initial"},
{streamgate.ObservationKindRecoveryDispatched, segRebound, uncommitted, "eval-initial"},
{streamgate.ObservationKindResponseStaged, segRebound, uncommitted, ""},
{streamgate.ObservationKindFilterEvaluationStarted, segRebound, uncommitted, "eval-rebound"},
{streamgate.ObservationKindFilterEvaluated, segRebound, uncommitted, "eval-rebound"},
{streamgate.ObservationKindArbitrationDecided, segRebound, uncommitted, "eval-rebound"},
{streamgate.ObservationKindReleaseCommitted, segRebound, streamgate.CommitStateStreamOpen, "eval-rebound"},
{streamgate.ObservationKindTerminalCommitted, segRebound, streamgate.CommitStateTerminalCommitted, "eval-rebound"},
}
epochByGroup := make(map[string]uint64)
idx := 0
for r, want := range rows {
match := -1
for j := idx; j < len(observations); j++ {
if observations[j].Kind() == want.kind {
match = j
break
}
}
if match < 0 {
t.Fatalf("required row %d kind=%s not found in order after index %d; timeline=%v", r, want.kind, idx, observationKinds(observations))
}
obs := observations[match]
wantID, wantTarget := "attempt.pool-attempt-1", initialTarget
if want.segment == segRebound {
wantID, wantTarget = "attempt.pool-attempt-2", reboundTarget
}
if obs.AttemptID() != wantID || obs.AttemptTarget() != wantTarget {
t.Fatalf("required row %d kind=%s attempt/target=(%q,%+v), want (%q,%+v)", r, want.kind, obs.AttemptID(), obs.AttemptTarget(), wantID, wantTarget)
}
if obs.CommitState() != want.commit {
t.Fatalf("required row %d kind=%s commit=%s, want %s", r, want.kind, obs.CommitState(), want.commit)
}
if want.epochGroup == "" {
if obs.EpochID() != 0 {
t.Fatalf("required row %d kind=%s epoch=%d, want 0 (unbound)", r, want.kind, obs.EpochID())
}
} else {
if obs.EpochID() == 0 {
t.Fatalf("required row %d kind=%s epoch=0, want non-zero epoch for group %s", r, want.kind, want.epochGroup)
}
if prev, ok := epochByGroup[want.epochGroup]; ok {
if prev != obs.EpochID() {
t.Fatalf("required row %d kind=%s epoch=%d, want %d consistent with group %s", r, want.kind, obs.EpochID(), prev, want.epochGroup)
}
} else {
epochByGroup[want.epochGroup] = obs.EpochID()
}
}
idx = match + 1
}
}
func TestStreamGateProviderPoolObservationCorrelationTimeline(t *testing.T) {
rawSentinels := []string{"SECRET_PROMPT_CONTENT", "SECRET_OUTPUT_CONTENT", "SECRET_TOOL_ARGS", "SECRET_AUTH_TOKEN", "SECRET_PREPARER_INPUT"}
rawBody := []byte(`{"model":"client-model","messages":[{"role":"user","content":"SECRET_PROMPT_CONTENT SECRET_TOOL_ARGS SECRET_AUTH_TOKEN SECRET_PREPARER_INPUT"}],"stream":true}`)
service := newScriptedPoolRunService(
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathNormalized), runID: "pool-attempt-1",
provider: "prov-normalized", target: "served-normalized",
runEvents: bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "SECRET_OUTPUT_CONTENT"},
&iop.RunEvent{Type: "complete"},
),
},
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathTunnel), runID: "pool-attempt-2",
provider: "prov-tunnel", target: "served-tunnel", frames: poolChatSSEFrames("switched"),
},
)
var logOutput bytes.Buffer
encoderConfig := zap.NewProductionEncoderConfig()
encoderConfig.TimeKey = ""
logger := zap.New(zapcore.NewCore(zapcore.NewJSONEncoder(encoderConfig), zapcore.AddSync(&logOutput), zapcore.InfoLevel))
recorder := &recordingOpenAIObservationSink{delegate: newZapFilterObservationSink(logger)}
srv, dc := buildStreamGatePoolChatFixture(t, service, rawBody, 3)
srv.SetObservationSink(recorder)
w, releaseSink, runErr := runPoolChatStreamGate(t, srv, dc, newInjectedViolationRegistration(t, dc, "test.injected_violation", 1))
if runErr != nil {
t.Fatalf("rt.Run: %v", runErr)
}
if !strings.Contains(w.body.String(), `"content":"switched"`) || strings.Contains(w.body.String(), "SECRET_OUTPUT_CONTENT") {
t.Fatalf("unexpected response body: %q", w.body.String())
}
composite := releaseSink.(*openAICompositeReleaseSink)
if committed, success := composite.terminalStatus(); !committed || !success {
t.Fatalf("terminal status=(%v,%v), want successful commit", committed, success)
}
observations := recorder.Snapshot()
assertProviderPoolObservationTimeline(t, observations)
eventDump, encoded := fmt.Sprintf("%+v", observations), logOutput.String()
for _, sentinel := range rawSentinels {
if strings.Contains(eventDump, sentinel) {
t.Fatalf("observation event contains raw sentinel %q", sentinel)
}
if strings.Contains(encoded, sentinel) {
t.Fatalf("observation log contains raw sentinel %q", sentinel)
}
}
for _, field := range []string{`"model_group":"client-model"`, `"actual_model":"served-normalized"`, `"actual_model":"served-tunnel"`, `"actual_provider":"prov-tunnel"`, `"execution_path":"provider_tunnel"`} {
if !strings.Contains(encoded, field) {
t.Errorf("observation log missing %s", field)
}
}
}
func TestStreamGateProviderPoolObservationSinkFailureIsolation(t *testing.T) {
rawBody := []byte(`{"model":"client-model","messages":[{"role":"user","content":"hi"}],"stream":true}`)
service := newScriptedPoolRunService(
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathNormalized), runID: "pool-attempt-1",
provider: "prov-normalized", target: "served-normalized",
runEvents: bufferedRunEvents(&iop.RunEvent{Type: "delta", Delta: "discarded"}, &iop.RunEvent{Type: "complete"}),
},
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathTunnel), runID: "pool-attempt-2",
provider: "prov-tunnel", target: "served-tunnel", frames: poolChatSSEFrames("switched"),
},
)
recorder := &recordingOpenAIObservationSink{fail: errors.New("injected observation sink failure")}
srv, dc := buildStreamGatePoolChatFixture(t, service, rawBody, 3)
srv.SetObservationSink(recorder)
w, releaseSink, runErr := runPoolChatStreamGate(t, srv, dc, newInjectedViolationRegistration(t, dc, "test.injected_violation", 1))
if runErr != nil {
t.Fatalf("sink failure escaped runtime: %v", runErr)
}
if !strings.Contains(w.body.String(), `"content":"switched"`) {
t.Fatalf("successful response missing after sink failures: %q", w.body.String())
}
composite := releaseSink.(*openAICompositeReleaseSink)
if committed, success := composite.terminalStatus(); !committed || !success {
t.Fatalf("terminal status=(%v,%v), want successful commit", committed, success)
}
if service.poolSubmits() != 2 {
t.Fatalf("pool submits=%d, want 2", service.poolSubmits())
}
_, cancels, _, _, _, _ := service.snapshot()
if n := countString(cancels, "pool-attempt-1"); n != 1 {
t.Fatalf("replaced attempt canceled %d times, want exactly once", n)
}
if n := countString(cancels, "pool-attempt-2"); n > 1 {
t.Fatalf("terminal attempt canceled %d times, want at most once", n)
}
assertProviderPoolObservationTimeline(t, recorder.Snapshot())
}
func countString(values []string, want string) int {
n := 0
for _, value := range values {
if value == want {
n++
}
}
return n
}
// TestStreamGateProviderPoolRecoverySwitchesActualPathOnce is the S16
// vertical-slice evidence: one RequestRuntime admits the initial provider-pool
// attempt, a blocking violation triggers a single re-admission through
// SubmitProviderPool, the pool selects a candidate on the *other* execution
// path, and only that attempt's framing and bytes reach the caller.
func TestStreamGateProviderPoolRecoverySwitchesActualPathOnce(t *testing.T) {
rawBody := []byte(`{"model":"client-model","messages":[{"role":"user","content":"hi"}],"stream":true}`)
t.Run("normalized to tunnel", func(t *testing.T) {
service := newScriptedPoolRunService(
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathNormalized), runID: "pool-attempt-1",
provider: "prov-normalized", target: "served-normalized",
runEvents: bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "discarded-normalized"},
&iop.RunEvent{Type: "complete"},
),
},
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathTunnel), runID: "pool-attempt-2",
provider: "prov-tunnel", target: "served-tunnel",
frames: poolChatSSEFrames("switched"),
},
)
srv, dc := buildStreamGatePoolChatFixture(t, service, rawBody, 3)
w, sink, runErr := runPoolChatStreamGate(t, srv, dc, newInjectedViolationRegistration(t, dc, "test.injected_violation", 1))
if runErr != nil {
t.Fatalf("rt.Run: %v", runErr)
}
if got := service.poolSubmits(); got != 2 {
t.Fatalf("SubmitProviderPool calls: got %d, want 2 (initial admission + exactly one re-admission)", got)
}
composite, ok := sink.(*openAICompositeReleaseSink)
if !ok {
t.Fatalf("provider-pool request must use the composite sink, got %T", sink)
}
if got := composite.resolvedCodec(); got != openAIStreamGateCodecTunnel {
t.Fatalf("resolved codec: got %q, want tunnel framing from the recovered attempt", got)
}
if got := w.headerCallCount(); got != 1 {
t.Fatalf("WriteHeader call count: got %d, want 1", got)
}
body := w.body.String()
if !strings.Contains(body, `"content":"switched"`) {
t.Fatalf("recovered tunnel attempt bytes missing: body=%q", body)
}
if strings.Contains(body, "discarded-normalized") {
t.Fatalf("discarded normalized attempt leaked into the response: body=%q", body)
}
if strings.Contains(body, "chat.completion.chunk") {
t.Fatalf("normalized SSE framing leaked into a tunnel-framed response: body=%q", body)
}
_, cancels, _, _, _, _ := service.snapshot()
if !containsString(cancels, "pool-attempt-1") {
t.Fatalf("the replaced attempt must be aborted, cancels=%v", cancels)
}
})
t.Run("tunnel to normalized", func(t *testing.T) {
service := newScriptedPoolRunService(
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathTunnel), runID: "pool-attempt-1",
provider: "prov-tunnel", target: "served-tunnel",
frames: poolChatSSEFrames("discarded-tunnel"),
},
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathNormalized), runID: "pool-attempt-2",
provider: "prov-normalized", target: "served-normalized",
runEvents: bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "switched back"},
&iop.RunEvent{Type: "complete"},
),
},
)
srv, dc := buildStreamGatePoolChatFixture(t, service, rawBody, 3)
w, sink, runErr := runPoolChatStreamGate(t, srv, dc, newInjectedViolationRegistration(t, dc, "test.injected_violation", 1))
if runErr != nil {
t.Fatalf("rt.Run: %v", runErr)
}
if got := service.poolSubmits(); got != 2 {
t.Fatalf("SubmitProviderPool calls: got %d, want 2", got)
}
composite := sink.(*openAICompositeReleaseSink)
if got := composite.resolvedCodec(); got != openAIStreamGateCodecNormalized {
t.Fatalf("resolved codec: got %q, want normalized framing from the recovered attempt", got)
}
if got := w.headerCallCount(); got != 1 {
t.Fatalf("WriteHeader call count: got %d, want 1", got)
}
chunks := parseSSEChatChunks(t, w.body.String())
if got := joinedContent(chunks); got != "switched back" {
t.Fatalf("joined content: got %q, want only the recovered normalized attempt's content", got)
}
if strings.Contains(w.body.String(), "discarded-tunnel") {
t.Fatalf("discarded tunnel attempt leaked into the response: body=%q", w.body.String())
}
})
}
// TestStreamGateProviderPoolExhaustionTerminatesWithoutPathLeak verifies the
// bounded budget still holds across a path switch: with a request-total cap of
// 1 the second violation terminates instead of re-admitting a third candidate,
// and no attempt's content is released.
func TestStreamGateProviderPoolExhaustionTerminatesWithoutPathLeak(t *testing.T) {
rawBody := []byte(`{"model":"client-model","messages":[{"role":"user","content":"hi"}],"stream":true}`)
service := newScriptedPoolRunService(
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathNormalized), runID: "pool-attempt-1",
provider: "prov-normalized", target: "served-normalized",
runEvents: bufferedRunEvents(&iop.RunEvent{Type: "delta", Delta: "still-bad-1"}, &iop.RunEvent{Type: "complete"}),
},
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathTunnel), runID: "pool-attempt-2",
provider: "prov-tunnel", target: "served-tunnel",
frames: poolChatSSEFrames("still-bad-2"),
},
)
srv, dc := buildStreamGatePoolChatFixture(t, service, rawBody, 1)
w, _, runErr := runPoolChatStreamGate(t, srv, dc, newInjectedViolationRegistration(t, dc, "test.injected_violation", 2))
if runErr != nil {
t.Fatalf("rt.Run: %v", runErr)
}
if got := service.poolSubmits(); got != 2 {
t.Fatalf("SubmitProviderPool calls: got %d, want 2 (cap=1 allows exactly one re-admission)", got)
}
body := w.body.String()
if !strings.Contains(body, `"error"`) {
t.Fatalf("expected an error terminal once the recovery budget is exhausted, body=%q", body)
}
for _, leak := range []string{"still-bad-1", "still-bad-2"} {
if strings.Contains(body, leak) {
t.Fatalf("exhausted recovery leaked attempt content %q: body=%q", leak, body)
}
}
}
// TestStreamGateProviderPoolSimultaneousViolationsDispatchOnce verifies that two
// independent blocking filters violating on the same terminal produce exactly
// one recovery plan and one re-admission, not one per filter.
func TestStreamGateProviderPoolSimultaneousViolationsDispatchOnce(t *testing.T) {
rawBody := []byte(`{"model":"client-model","messages":[{"role":"user","content":"hi"}],"stream":true}`)
service := newScriptedPoolRunService(
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathNormalized), runID: "pool-attempt-1",
provider: "prov-normalized", target: "served-normalized",
runEvents: bufferedRunEvents(&iop.RunEvent{Type: "delta", Delta: "discarded"}, &iop.RunEvent{Type: "complete"}),
},
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathNormalized), runID: "pool-attempt-2",
provider: "prov-normalized", target: "served-normalized",
runEvents: bufferedRunEvents(&iop.RunEvent{Type: "delta", Delta: "final answer"}, &iop.RunEvent{Type: "complete"}),
},
)
srv, dc := buildStreamGatePoolChatFixture(t, service, rawBody, 3)
w, _, runErr := runPoolChatStreamGate(t, srv, dc,
newInjectedViolationRegistration(t, dc, "test.injected_violation_a", 1),
newInjectedViolationRegistration(t, dc, "test.injected_violation_b", 1),
)
if runErr != nil {
t.Fatalf("rt.Run: %v", runErr)
}
if got := service.poolSubmits(); got != 2 {
t.Fatalf("SubmitProviderPool calls: got %d, want 2 (two simultaneous violations must share one plan/dispatch)", got)
}
chunks := parseSSEChatChunks(t, w.body.String())
if got := joinedContent(chunks); got != "final answer" {
t.Fatalf("joined content: got %q, want the single recovered attempt's content", got)
}
if got := w.headerCallCount(); got != 1 {
t.Fatalf("WriteHeader call count: got %d, want 1", got)
}
}
type devRepeatGuardSmokeEvidence struct {
Run int `json:"run"`
Attempt int `json:"attempt"`
Model string `json:"model"`
StatusCode int `json:"status_code"`
Terminal string `json:"terminal"`
ProviderResponseIDs []string `json:"provider_response_ids"`
downstreamRepeated bool
}
type devRepeatGuardLifecycleEvidence struct {
Run int `json:"run"`
Model string `json:"model"`
ActualModel string `json:"actual_model"`
Provider string `json:"provider"`
CorrelationID string `json:"correlation_id"`
Result string `json:"result"`
Decision string `json:"decision"`
Lifecycle string `json:"lifecycle"`
TerminalReason string `json:"terminal_reason"`
Fingerprint string `json:"fingerprint,omitempty"`
Offset int `json:"offset"`
}
type devRepeatGuardSmokeSummary struct {
Result string `json:"result"`
Model string `json:"model"`
Concurrency int `json:"concurrency"`
Runs int `json:"runs"`
ProviderResponses []devRepeatGuardSmokeEvidence `json:"provider_responses"`
Lifecycles []devRepeatGuardLifecycleEvidence `json:"lifecycles"`
CapacityRuns []devRepeatGuardCapacityEvidence `json:"capacity_runs"`
}
type devRepeatGuardCapacityEvidence struct {
Run int `json:"run"`
ConfiguredCapacity int `json:"configured_capacity"`
PeakInFlight int `json:"peak_in_flight"`
PeakQueued int `json:"peak_queued"`
FinalInFlight int `json:"final_in_flight"`
FinalQueued int `json:"final_queued"`
}
type devRepeatGuardCapacitySnapshot struct {
configuredCapacity int
inFlight int
queued int
}
type devRepeatGuardObservation struct {
Message string `json:"msg"`
Sequence uint64 `json:"sequence"`
ObservationKind string `json:"observation_kind"`
CorrelationID string `json:"correlation_id"`
AttemptID string `json:"attempt_id"`
ModelGroup string `json:"model_group"`
ActualModel string `json:"actual_model"`
ActualProvider string `json:"actual_provider"`
FilterID string `json:"filter_id"`
RuleID string `json:"rule_id"`
FilterOutcome string `json:"filter_outcome"`
DecisionKind string `json:"decision_kind"`
ArbitrationAction string `json:"arbitration_action"`
TerminalReason string `json:"terminal_reason"`
RecoveryStrategy string `json:"recovery_strategy"`
EvidenceDescriptorCode string `json:"evidence_descriptor_code"`
EvidenceFingerprint string `json:"evidence_fingerprint"`
EvidenceOffset int `json:"evidence_offset"`
line int
}
type devRepeatGuardSSEEvidence struct {
terminal string
providerResponseIDs []string
}
const (
devRepeatGuardMaxProviderResponseIDs = 4
devRepeatGuardDefaultRequestTimeout = 25 * time.Minute
devRepeatGuardMaxRequestTimeout = 25 * time.Minute
devRepeatGuardPerRunSettlementAllowance = 2 * time.Minute
devRepeatGuardFinalCleanupReserve = 5 * time.Minute
devRepeatGuardRequestTimeoutEnvironmentName = "IOP_OPENAI_SMOKE_REQUEST_TIMEOUT"
)
// TestDevRepeatGuardOrnithSmoke is an env-gated live entry. Raw prompt and SSE
// bytes are read/written only in ignored agent-test/runs paths. Stdout and the
// sanitized summary contain no prompt, output, token, tool, auth, or endpoint
// value. Reproduction is classified only from correlated raw-free guard
// observations; released text is inspected only to fail a run if a lifecycle
// observation says a blocked duplicate reached the downstream stream.
func TestDevRepeatGuardOrnithSmoke(t *testing.T) {
baseURL := strings.TrimRight(strings.TrimSpace(os.Getenv("IOP_OPENAI_BASE_URL")), "/")
model := strings.TrimSpace(os.Getenv("IOP_OPENAI_MODEL"))
promptFile := strings.TrimSpace(os.Getenv("IOP_OPENAI_SMOKE_PROMPT_FILE"))
artifactDir := strings.TrimSpace(os.Getenv("IOP_OPENAI_SMOKE_ARTIFACT_DIR"))
observationFile := strings.TrimSpace(os.Getenv("IOP_OPENAI_SMOKE_OBSERVATION_FILE"))
statusURL := strings.TrimSpace(os.Getenv("IOP_OPENAI_SMOKE_STATUS_URL"))
rawProviderIDs := strings.TrimSpace(os.Getenv("IOP_OPENAI_SMOKE_PROVIDER_IDS"))
required := map[string]string{
"IOP_OPENAI_BASE_URL": baseURL,
"IOP_OPENAI_MODEL": model,
"IOP_OPENAI_SMOKE_PROMPT_FILE": promptFile,
"IOP_OPENAI_SMOKE_ARTIFACT_DIR": artifactDir,
"IOP_OPENAI_SMOKE_OBSERVATION_FILE": observationFile,
"IOP_OPENAI_SMOKE_STATUS_URL": statusURL,
"IOP_OPENAI_SMOKE_PROVIDER_IDS": rawProviderIDs,
}
configured := 0
for _, value := range required {
if value != "" {
configured++
}
}
if configured == 0 {
t.Skip("live repeat-guard smoke environment is not configured")
}
for key, value := range required {
if value == "" {
t.Fatalf("%s is required when the live repeat-guard smoke is configured", key)
}
}
providerIDs, err := parseDevRepeatGuardProviderIDs(rawProviderIDs)
if err != nil {
t.Fatalf("parse selected provider IDs: %v", err)
}
concurrency := devRepeatGuardEnvInt(t, "IOP_OPENAI_SMOKE_CONCURRENCY", 5)
runs := devRepeatGuardEnvInt(t, "IOP_OPENAI_SMOKE_RUNS", 3)
if concurrency != 5 || runs != 3 {
t.Fatalf("live repeat-guard evidence requires concurrency=5 and runs=3")
}
requestTimeout := devRepeatGuardRequestTimeout(t, runs)
prompt, err := os.ReadFile(promptFile)
if err != nil {
t.Fatalf("read ignored smoke prompt: %v", err)
}
if len(prompt) == 0 {
t.Fatal("ignored smoke prompt is empty")
}
if err := os.MkdirAll(artifactDir, 0o700); err != nil {
t.Fatalf("create ignored smoke artifact directory: %v", err)
}
client := &http.Client{Timeout: requestTimeout}
statusClient := &http.Client{Timeout: 5 * time.Second}
providerResponses := make([]devRepeatGuardSmokeEvidence, 0, concurrency*runs)
lifecycles := make([]devRepeatGuardLifecycleEvidence, 0, concurrency*runs)
capacityRuns := make([]devRepeatGuardCapacityEvidence, 0, runs)
for run := 1; run <= runs; run++ {
if _, err := waitDevRepeatGuardCapacityIdle(
statusClient, statusURL, providerIDs, 30*time.Second,
); err != nil {
t.Fatalf("wait for idle selected capacity before run=%d: %v", run, err)
}
observationOffset, err := devRepeatGuardObservationQuietOffset(
observationFile, 2*time.Second,
)
if err != nil {
t.Fatalf("observation sink is not quiescent before run=%d: %v", run, err)
}
var (
wg sync.WaitGroup
resultC = make(chan devRepeatGuardSmokeEvidence, concurrency)
errorC = make(chan error, concurrency)
startC = make(chan struct{})
doneC = make(chan struct{})
)
for attempt := 1; attempt <= concurrency; attempt++ {
wg.Add(1)
go func(run, attempt int) {
defer wg.Done()
<-startC
evidence, requestErr := runDevRepeatGuardSmokeAttempt(
client, baseURL, model, string(prompt), artifactDir, run, attempt,
)
if requestErr != nil {
errorC <- requestErr
return
}
resultC <- evidence
}(run, attempt)
}
go func() {
wg.Wait()
close(doneC)
}()
capacityC := make(chan devRepeatGuardCapacityEvidence, 1)
capacityErrC := make(chan error, 1)
observerReadyC := make(chan struct{})
go func(run int) {
evidence, observeErr := observeDevRepeatGuardCapacityRun(
statusClient, statusURL, providerIDs, run, doneC, observerReadyC,
requestTimeout+devRepeatGuardPerRunSettlementAllowance,
)
if observeErr != nil {
capacityErrC <- observeErr
return
}
capacityC <- evidence
}(run)
<-observerReadyC
close(startC)
wg.Wait()
close(resultC)
close(errorC)
for requestErr := range errorC {
if requestErr != nil {
t.Errorf("live repeat-guard request failed: %v", requestErr)
}
}
batch := make([]devRepeatGuardSmokeEvidence, 0, concurrency)
for evidence := range resultC {
batch = append(batch, evidence)
}
if t.Failed() {
return
}
select {
case capacity := <-capacityC:
capacityRuns = append(capacityRuns, capacity)
case capacityErr := <-capacityErrC:
t.Fatalf("observe selected provider capacity for run=%d: %v", run, capacityErr)
}
sort.Slice(batch, func(i, j int) bool {
return batch[i].Attempt < batch[j].Attempt
})
observations, err := waitDevRepeatGuardObservationSegment(
observationFile, observationOffset, model, concurrency, prompt,
[]byte(strings.TrimSpace(os.Getenv("IOP_OPENAI_API_KEY"))),
)
if err != nil {
t.Fatalf("read observation segment for run=%d: %v", run, err)
}
correlated, err := correlateDevRepeatGuardBatch(batch, observations)
if err != nil {
t.Fatalf("correlate observation segment for run=%d: %v", run, err)
}
providerResponses = append(providerResponses, batch...)
lifecycles = append(lifecycles, correlated...)
}
if t.Failed() {
return
}
if len(providerResponses) != 15 || len(lifecycles) != 15 || len(capacityRuns) != 3 {
t.Fatalf(
"sanitized evidence rows: provider_responses=%d lifecycles=%d capacity_runs=%d, want 15, 15, and 3",
len(providerResponses), len(lifecycles), len(capacityRuns),
)
}
sort.Slice(providerResponses, func(i, j int) bool {
if providerResponses[i].Run == providerResponses[j].Run {
return providerResponses[i].Attempt < providerResponses[j].Attempt
}
return providerResponses[i].Run < providerResponses[j].Run
})
sort.Slice(lifecycles, func(i, j int) bool {
if lifecycles[i].Run == lifecycles[j].Run {
return lifecycles[i].CorrelationID < lifecycles[j].CorrelationID
}
return lifecycles[i].Run < lifecycles[j].Run
})
if err := validateDevRepeatGuardUniqueEvidence(providerResponses, lifecycles); err != nil {
t.Fatalf("validate final evidence identity: %v", err)
}
overall := "not_reproduced"
for _, result := range lifecycles {
if result.Result == "reproduced" {
overall = "reproduced"
break
}
}
summary := devRepeatGuardSmokeSummary{
Result: overall, Model: model, Concurrency: concurrency, Runs: runs,
ProviderResponses: providerResponses, Lifecycles: lifecycles, CapacityRuns: capacityRuns,
}
encoded, err := json.MarshalIndent(summary, "", " ")
if err != nil {
t.Fatalf("encode sanitized smoke summary: %v", err)
}
encoded = append(encoded, '\n')
if err := os.WriteFile(filepath.Join(artifactDir, "summary.json"), encoded, 0o600); err != nil {
t.Fatalf("write sanitized smoke summary: %v", err)
}
t.Logf(
"repeat-guard live result=%s provider_responses=%d edge_correlations=%d capacity_runs=%d",
overall, len(providerResponses), len(lifecycles), len(capacityRuns),
)
}
func validateDevRepeatGuardUniqueEvidence(
providerResponses []devRepeatGuardSmokeEvidence,
lifecycles []devRepeatGuardLifecycleEvidence,
) error {
providerIDs := make(map[string]struct{}, len(providerResponses))
for _, response := range providerResponses {
if err := collectDevRepeatGuardProviderResponseIDs(
response.ProviderResponseIDs, providerIDs,
fmt.Sprintf("run=%d attempt=%d", response.Run, response.Attempt),
); err != nil {
return err
}
}
correlations := make(map[string]struct{}, len(lifecycles))
for _, lifecycle := range lifecycles {
correlation := strings.TrimSpace(lifecycle.CorrelationID)
if correlation == "" {
return errors.New("Edge lifecycle correlation is empty")
}
if _, found := correlations[correlation]; found {
return fmt.Errorf("Edge lifecycle correlation %q is not globally unique", correlation)
}
correlations[correlation] = struct{}{}
}
return nil
}
func devRepeatGuardRequestTimeout(t *testing.T, runs int) time.Duration {
t.Helper()
deadline, hasDeadline := t.Deadline()
requestTimeout, err := resolveDevRepeatGuardRequestTimeout(
os.Getenv(devRepeatGuardRequestTimeoutEnvironmentName),
runs,
time.Now(),
deadline,
hasDeadline,
)
if err != nil {
t.Fatalf("%s: %v", devRepeatGuardRequestTimeoutEnvironmentName, err)
}
return requestTimeout
}
func resolveDevRepeatGuardRequestTimeout(
raw string,
runs int,
now time.Time,
deadline time.Time,
hasDeadline bool,
) (time.Duration, error) {
if runs <= 0 || runs > 32 {
return 0, errors.New("run count must be between 1 and 32")
}
requestTimeout := devRepeatGuardDefaultRequestTimeout
if value := strings.TrimSpace(raw); value != "" {
parsed, err := time.ParseDuration(value)
if err != nil {
return 0, errors.New("request timeout must be a valid Go duration")
}
requestTimeout = parsed
}
if requestTimeout <= 0 {
return 0, errors.New("request timeout must be positive")
}
if requestTimeout > devRepeatGuardMaxRequestTimeout {
return 0, fmt.Errorf(
"request timeout must not exceed %s",
devRepeatGuardMaxRequestTimeout,
)
}
requiredBudget := time.Duration(runs) *
(requestTimeout + devRepeatGuardPerRunSettlementAllowance)
requiredBudget += devRepeatGuardFinalCleanupReserve
if hasDeadline && deadline.Sub(now) < requiredBudget {
return 0, fmt.Errorf(
"test deadline leaves %s, need at least %s for %d runs, settlement, and cleanup",
deadline.Sub(now), requiredBudget, runs,
)
}
return requestTimeout, nil
}
func collectDevRepeatGuardProviderResponseIDs(
ids []string,
seen map[string]struct{},
context string,
) error {
if len(ids) == 0 || len(ids) > devRepeatGuardMaxProviderResponseIDs {
return fmt.Errorf(
"%s provider response identity count=%d, want 1..%d",
context, len(ids), devRepeatGuardMaxProviderResponseIDs,
)
}
const chatCompletionPrefix = "chatcmpl-"
local := make(map[string]struct{}, len(ids))
for _, rawID := range ids {
id := strings.TrimSpace(rawID)
if !strings.HasPrefix(id, chatCompletionPrefix) ||
strings.TrimPrefix(id, chatCompletionPrefix) == "" {
return fmt.Errorf("%s has invalid provider response identity", context)
}
if _, found := local[id]; found {
return fmt.Errorf("%s repeats provider response identity %q", context, id)
}
local[id] = struct{}{}
if _, found := seen[id]; found {
return fmt.Errorf("provider response identity %q is not globally unique", id)
}
seen[id] = struct{}{}
}
return nil
}
func devRepeatGuardEnvInt(t *testing.T, key string, fallback int) int {
t.Helper()
raw := strings.TrimSpace(os.Getenv(key))
if raw == "" {
return fallback
}
value, err := strconv.Atoi(raw)
if err != nil || value <= 0 || value > 32 {
t.Fatalf("%s must be between 1 and 32", key)
}
return value
}
func parseDevRepeatGuardProviderIDs(raw string) ([]string, error) {
parts := strings.Split(raw, ",")
if len(parts) == 0 || len(parts) > 16 {
return nil, errors.New("provider ID count must be between 1 and 16")
}
seen := make(map[string]struct{}, len(parts))
providerIDs := make([]string, 0, len(parts))
for _, part := range parts {
providerID := strings.TrimSpace(part)
if providerID == "" {
return nil, errors.New("provider IDs must be non-empty")
}
if _, found := seen[providerID]; found {
return nil, fmt.Errorf("duplicate provider ID %q", providerID)
}
seen[providerID] = struct{}{}
providerIDs = append(providerIDs, providerID)
}
return providerIDs, nil
}
func parseDevRepeatGuardCapacityStatus(
raw []byte,
providerIDs []string,
) (devRepeatGuardCapacitySnapshot, error) {
var status struct {
Nodes []struct {
ProviderSnapshots []struct {
ID string `json:"id"`
Capacity int `json:"capacity"`
InFlight int `json:"in_flight"`
Queued int `json:"queued"`
} `json:"provider_snapshots"`
} `json:"nodes"`
}
if err := json.Unmarshal(raw, &status); err != nil {
return devRepeatGuardCapacitySnapshot{}, errors.New("status response is not valid JSON")
}
selected := make(map[string]struct{}, len(providerIDs))
for _, providerID := range providerIDs {
if strings.TrimSpace(providerID) == "" {
return devRepeatGuardCapacitySnapshot{}, errors.New("selected provider ID is empty")
}
selected[providerID] = struct{}{}
}
found := make(map[string]struct{}, len(providerIDs))
var snapshot devRepeatGuardCapacitySnapshot
for _, node := range status.Nodes {
for _, provider := range node.ProviderSnapshots {
if _, wanted := selected[provider.ID]; !wanted {
continue
}
if _, duplicate := found[provider.ID]; duplicate {
return devRepeatGuardCapacitySnapshot{}, fmt.Errorf(
"selected provider %q appears more than once", provider.ID,
)
}
if provider.Capacity < 0 || provider.InFlight < 0 || provider.Queued < 0 {
return devRepeatGuardCapacitySnapshot{}, fmt.Errorf(
"selected provider %q has a negative capacity counter", provider.ID,
)
}
found[provider.ID] = struct{}{}
snapshot.configuredCapacity += provider.Capacity
snapshot.inFlight += provider.InFlight
snapshot.queued += provider.Queued
}
}
for _, providerID := range providerIDs {
if _, ok := found[providerID]; !ok {
return devRepeatGuardCapacitySnapshot{}, fmt.Errorf(
"selected provider %q is missing from status", providerID,
)
}
}
return snapshot, nil
}
func fetchDevRepeatGuardCapacityStatus(
client *http.Client,
statusURL string,
providerIDs []string,
) (devRepeatGuardCapacitySnapshot, error) {
req, err := http.NewRequest(http.MethodGet, statusURL, nil)
if err != nil {
return devRepeatGuardCapacitySnapshot{}, errors.New("build status request")
}
resp, err := client.Do(req)
if err != nil {
return devRepeatGuardCapacitySnapshot{}, errors.New("status request failed")
}
defer resp.Body.Close()
if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
return devRepeatGuardCapacitySnapshot{}, fmt.Errorf(
"status request returned HTTP %d", resp.StatusCode,
)
}
const statusBodyLimit = 4 * 1024 * 1024
raw, err := io.ReadAll(io.LimitReader(resp.Body, statusBodyLimit+1))
if err != nil {
return devRepeatGuardCapacitySnapshot{}, errors.New("read status response")
}
if len(raw) > statusBodyLimit {
return devRepeatGuardCapacitySnapshot{}, errors.New("status response exceeds 4 MiB")
}
return parseDevRepeatGuardCapacityStatus(raw, providerIDs)
}
func waitDevRepeatGuardCapacityIdle(
client *http.Client,
statusURL string,
providerIDs []string,
timeout time.Duration,
) (devRepeatGuardCapacitySnapshot, error) {
deadline := time.Now().Add(timeout)
var lastErr error
for {
snapshot, err := fetchDevRepeatGuardCapacityStatus(client, statusURL, providerIDs)
if err == nil {
if snapshot.configuredCapacity != 4 {
return devRepeatGuardCapacitySnapshot{}, fmt.Errorf(
"selected provider capacity=%d, want 4", snapshot.configuredCapacity,
)
}
if snapshot.inFlight == 0 && snapshot.queued == 0 {
return snapshot, nil
}
lastErr = fmt.Errorf(
"selected providers are not idle: in_flight=%d queued=%d",
snapshot.inFlight, snapshot.queued,
)
} else {
lastErr = err
}
if time.Now().After(deadline) {
return devRepeatGuardCapacitySnapshot{}, lastErr
}
time.Sleep(100 * time.Millisecond)
}
}
func observeDevRepeatGuardCapacityRun(
client *http.Client,
statusURL string,
providerIDs []string,
run int,
requestsDone <-chan struct{},
observerReady chan<- struct{},
timeout time.Duration,
) (devRepeatGuardCapacityEvidence, error) {
close(observerReady)
deadline := time.Now().Add(timeout)
requestsComplete := false
evidence := devRepeatGuardCapacityEvidence{Run: run}
for {
snapshot, err := fetchDevRepeatGuardCapacityStatus(client, statusURL, providerIDs)
if err != nil {
return devRepeatGuardCapacityEvidence{}, err
}
if snapshot.configuredCapacity != 4 {
return devRepeatGuardCapacityEvidence{}, fmt.Errorf(
"selected provider capacity=%d, want 4", snapshot.configuredCapacity,
)
}
evidence.ConfiguredCapacity = snapshot.configuredCapacity
if snapshot.inFlight > evidence.PeakInFlight {
evidence.PeakInFlight = snapshot.inFlight
}
if snapshot.queued > evidence.PeakQueued {
evidence.PeakQueued = snapshot.queued
}
if !requestsComplete {
select {
case <-requestsDone:
requestsComplete = true
default:
}
}
if requestsComplete && snapshot.inFlight == 0 && snapshot.queued == 0 {
evidence.FinalInFlight = snapshot.inFlight
evidence.FinalQueued = snapshot.queued
if err := validateDevRepeatGuardCapacityEvidence(evidence); err != nil {
return devRepeatGuardCapacityEvidence{}, err
}
return evidence, nil
}
if time.Now().After(deadline) {
return devRepeatGuardCapacityEvidence{}, errors.New(
"selected provider capacity did not recover before the deadline",
)
}
time.Sleep(100 * time.Millisecond)
}
}
func validateDevRepeatGuardCapacityEvidence(evidence devRepeatGuardCapacityEvidence) error {
if evidence.Run <= 0 {
return errors.New("capacity evidence has no run index")
}
if evidence.ConfiguredCapacity != 4 {
return fmt.Errorf(
"configured capacity=%d, want 4", evidence.ConfiguredCapacity,
)
}
if evidence.PeakInFlight != 4 {
return fmt.Errorf("peak in_flight=%d, want 4", evidence.PeakInFlight)
}
if evidence.PeakQueued < 1 {
return fmt.Errorf("peak queued=%d, want at least 1", evidence.PeakQueued)
}
if evidence.FinalInFlight != 0 || evidence.FinalQueued != 0 {
return fmt.Errorf(
"final in_flight/queued=%d/%d, want 0/0",
evidence.FinalInFlight, evidence.FinalQueued,
)
}
return nil
}
func runDevRepeatGuardSmokeAttempt(
client *http.Client,
baseURL, model, prompt, artifactDir string,
run, attempt int,
) (devRepeatGuardSmokeEvidence, error) {
body, err := json.Marshal(map[string]any{
"model": model,
"messages": []any{
map[string]any{"role": "user", "content": prompt},
},
"stream": true,
})
if err != nil {
return devRepeatGuardSmokeEvidence{}, fmt.Errorf("encode request: %w", err)
}
req, err := http.NewRequest(http.MethodPost, baseURL+"/v1/chat/completions", bytes.NewReader(body))
if err != nil {
return devRepeatGuardSmokeEvidence{}, fmt.Errorf("build request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
if apiKey := strings.TrimSpace(os.Getenv("IOP_OPENAI_API_KEY")); apiKey != "" {
req.Header.Set("Authorization", "Bearer "+apiKey)
}
resp, err := client.Do(req)
if err != nil {
return devRepeatGuardSmokeEvidence{}, fmt.Errorf("dispatch run=%d attempt=%d: %w", run, attempt, err)
}
defer resp.Body.Close()
raw, err := io.ReadAll(io.LimitReader(resp.Body, 64*1024*1024))
if err != nil {
return devRepeatGuardSmokeEvidence{}, fmt.Errorf("read run=%d attempt=%d: %w", run, attempt, err)
}
rawPath := filepath.Join(artifactDir, fmt.Sprintf("run-%02d-attempt-%02d.sse", run, attempt))
if err := os.WriteFile(rawPath, raw, 0o600); err != nil {
return devRepeatGuardSmokeEvidence{}, fmt.Errorf("write ignored raw run artifact: %w", err)
}
if resp.StatusCode < http.StatusOK || resp.StatusCode >= http.StatusMultipleChoices {
return devRepeatGuardSmokeEvidence{}, fmt.Errorf(
"run=%d attempt=%d returned non-2xx status %d", run, attempt, resp.StatusCode,
)
}
sse, err := validateDevRepeatGuardSSE(raw)
if err != nil {
return devRepeatGuardSmokeEvidence{}, fmt.Errorf(
"run=%d attempt=%d invalid SSE: %w", run, attempt, err,
)
}
channels := devRepeatGuardSmokeChannels(raw)
downstreamRepeated := channels.contentRepeated(64) || channels.reasoningRepeated(64)
evidence := devRepeatGuardSmokeEvidence{
Run: run, Attempt: attempt, Model: model, StatusCode: resp.StatusCode,
Terminal: sse.terminal, ProviderResponseIDs: sse.providerResponseIDs,
downstreamRepeated: downstreamRepeated,
}
return evidence, nil
}
func validateDevRepeatGuardSSE(raw []byte) (devRepeatGuardSSEEvidence, error) {
normalized := bytes.ReplaceAll(raw, []byte("\r\n"), []byte("\n"))
lines := bytes.Split(normalized, []byte{'\n'})
var (
dataLines [][]byte
eventCount int
errorCount int
successFinishCount int
sawDone bool
currentResponseID string
providerResponseIDs []string
closedResponseIDs = make(map[string]struct{})
)
processCurrentData := func() error {
if len(dataLines) == 0 {
return nil
}
eventCount++
data := bytes.Join(dataLines, []byte{'\n'})
dataLines = dataLines[:0]
if bytes.Equal(bytes.TrimSpace(data), []byte("[DONE]")) {
if sawDone {
return errors.New("duplicate [DONE] terminal marker")
}
if (errorCount == 1) == (successFinishCount == 1) {
return fmt.Errorf(
"[DONE] has invalid logical terminal state: errors=%d success_finishes=%d",
errorCount, successFinishCount,
)
}
sawDone = true
return nil
}
if sawDone {
return fmt.Errorf("data event %d appears after [DONE]", eventCount)
}
var envelope map[string]json.RawMessage
if err := json.Unmarshal(data, &envelope); err != nil {
return fmt.Errorf("data event %d is not JSON", eventCount)
}
hasError := false
if rawError := envelope["error"]; len(rawError) > 0 &&
!bytes.Equal(bytes.TrimSpace(rawError), []byte("null")) {
hasError = true
}
rawID := envelope["id"]
if len(rawID) > 0 {
var eventResponseID string
if err := json.Unmarshal(rawID, &eventResponseID); err != nil {
return fmt.Errorf("data event %d has a non-string response id", eventCount)
}
eventResponseID = strings.TrimSpace(eventResponseID)
if eventResponseID == "" {
return fmt.Errorf("data event %d has an empty response id", eventCount)
}
const chatCompletionPrefix = "chatcmpl-"
if !strings.HasPrefix(eventResponseID, chatCompletionPrefix) ||
strings.TrimPrefix(eventResponseID, chatCompletionPrefix) == "" {
return fmt.Errorf(
"data event %d has an invalid provider-native response identity",
eventCount,
)
}
if eventResponseID != currentResponseID {
if _, reused := closedResponseIDs[eventResponseID]; reused {
return fmt.Errorf(
"data event %d reused a non-contiguous provider response identity",
eventCount,
)
}
if len(providerResponseIDs) == devRepeatGuardMaxProviderResponseIDs {
return fmt.Errorf(
"data event %d exceeds the provider response identity bound of %d",
eventCount, devRepeatGuardMaxProviderResponseIDs,
)
}
if currentResponseID != "" {
closedResponseIDs[currentResponseID] = struct{}{}
}
currentResponseID = eventResponseID
providerResponseIDs = append(providerResponseIDs, eventResponseID)
}
} else if !hasError {
return fmt.Errorf(
"data event %d has no provider-native response identity",
eventCount,
)
}
finishes := 0
if rawChoices := envelope["choices"]; len(rawChoices) > 0 &&
!bytes.Equal(bytes.TrimSpace(rawChoices), []byte("null")) {
var choices []struct {
FinishReason json.RawMessage `json:"finish_reason"`
}
if err := json.Unmarshal(rawChoices, &choices); err != nil {
return fmt.Errorf("data event %d has invalid choices", eventCount)
}
for _, choice := range choices {
rawFinish := bytes.TrimSpace(choice.FinishReason)
if len(rawFinish) == 0 || bytes.Equal(rawFinish, []byte("null")) {
continue
}
var finish string
if err := json.Unmarshal(rawFinish, &finish); err != nil ||
strings.TrimSpace(finish) == "" {
return fmt.Errorf("data event %d has invalid finish_reason", eventCount)
}
finishes++
}
}
switch {
case hasError && (successFinishCount > 0 || finishes > 0):
return errors.New("SSE mixes error and success finish terminals")
case finishes > 0 && errorCount > 0:
return errors.New("SSE mixes success finish and error terminals")
case hasError && errorCount > 0:
return errors.New("SSE has duplicate error terminals")
case finishes > 0 && successFinishCount > 0:
return errors.New("SSE has duplicate success finish terminals")
case (errorCount > 0 || successFinishCount > 0) && !hasError && finishes == 0:
return errors.New("SSE has data after its logical terminal")
}
if hasError {
errorCount++
}
successFinishCount += finishes
if successFinishCount > 1 {
return fmt.Errorf("success finish count=%d, want 1", successFinishCount)
}
return nil
}
for _, rawLine := range lines {
line := bytes.TrimSuffix(rawLine, []byte{'\r'})
if len(line) == 0 {
if err := processCurrentData(); err != nil {
return devRepeatGuardSSEEvidence{}, err
}
continue
}
if line[0] == ':' {
continue
}
field, value, found := bytes.Cut(line, []byte{':'})
if !found {
field, value = line, nil
} else if len(value) > 0 && value[0] == ' ' {
value = value[1:]
}
switch string(field) {
case "data":
dataLines = append(dataLines, append([]byte(nil), value...))
case "event", "id", "retry":
// Valid SSE fields that do not carry the OpenAI-compatible payload.
default:
return devRepeatGuardSSEEvidence{}, fmt.Errorf("unsupported SSE field %q", field)
}
}
if err := processCurrentData(); err != nil {
return devRepeatGuardSSEEvidence{}, err
}
if eventCount == 0 {
return devRepeatGuardSSEEvidence{}, errors.New("no SSE data events")
}
if !sawDone {
return devRepeatGuardSSEEvidence{}, errors.New("SSE is missing final [DONE]")
}
if (errorCount == 1) == (successFinishCount == 1) {
return devRepeatGuardSSEEvidence{}, fmt.Errorf(
"invalid logical terminal state: errors=%d success_finishes=%d",
errorCount, successFinishCount,
)
}
terminal := "done"
if errorCount == 1 {
terminal = "error"
}
evidence := devRepeatGuardSSEEvidence{terminal: terminal}
if len(providerResponseIDs) == 0 {
return devRepeatGuardSSEEvidence{}, errors.New(
"SSE response has no provider-native chat completion identity",
)
}
evidence.providerResponseIDs = append([]string(nil), providerResponseIDs...)
return evidence, nil
}
// devRepeatGuardSmokeChannelEvidence carries the downstream content and
// reasoning channels as independent rune streams. Repetition must be evaluated
// per channel, never across the channel boundary, otherwise identical tails in
// two distinct channels produce a false downstream repeat.
type devRepeatGuardSmokeChannelEvidence struct {
content []rune
reasoning []rune
}
// contentRepeated reports whether the content channel alone contains a
// duplicated suffix of at least minimum runes.
func (e devRepeatGuardSmokeChannelEvidence) contentRepeated(minimum int) bool {
_, _, repeated := openAIRepeatedSuffix(e.content, minimum)
return repeated
}
// reasoningRepeated reports whether the canonical reasoning channel alone
// contains a duplicated suffix of at least minimum runes.
func (e devRepeatGuardSmokeChannelEvidence) reasoningRepeated(minimum int) bool {
_, _, repeated := openAIRepeatedSuffix(e.reasoning, minimum)
return repeated
}
// devRepeatGuardSmokeChannels reconstructs the content and reasoning channels
// from a raw SSE stream. It mirrors decodeChatTunnelFrame reasoning alias
// precedence: reasoning_content wins over reasoning, and the two aliases are
// never concatenated for a single choice.
func devRepeatGuardSmokeChannels(raw []byte) devRepeatGuardSmokeChannelEvidence {
var content, reasoning strings.Builder
for _, line := range bytes.Split(raw, []byte{'\n'}) {
line = bytes.TrimSpace(line)
if !bytes.HasPrefix(line, []byte("data:")) {
continue
}
data := bytes.TrimSpace(bytes.TrimPrefix(line, []byte("data:")))
if bytes.Equal(data, []byte("[DONE]")) {
continue
}
var chunk struct {
Choices []struct {
Delta struct {
Content string `json:"content"`
ReasoningContent string `json:"reasoning_content"`
Reasoning string `json:"reasoning"`
} `json:"delta"`
} `json:"choices"`
}
if json.Unmarshal(data, &chunk) != nil {
continue
}
for _, choice := range chunk.Choices {
content.WriteString(choice.Delta.Content)
canonicalReasoning := choice.Delta.ReasoningContent
if canonicalReasoning == "" {
canonicalReasoning = choice.Delta.Reasoning
}
reasoning.WriteString(canonicalReasoning)
}
}
return devRepeatGuardSmokeChannelEvidence{
content: []rune(content.String()),
reasoning: []rune(reasoning.String()),
}
}
func devRepeatGuardObservationQuietOffset(path string, quietWindow time.Duration) (int64, error) {
if quietWindow < 0 {
return 0, errors.New("observation quiet window must not be negative")
}
return devRepeatGuardObservationQuietOffsetWithWait(path, func() error {
timer := time.NewTimer(quietWindow)
defer timer.Stop()
<-timer.C
return nil
})
}
func devRepeatGuardObservationQuietOffsetWithWait(
path string,
wait func() error,
) (int64, error) {
before, err := os.Stat(path)
if err != nil {
return 0, err
}
if !before.Mode().IsRegular() {
return 0, errors.New("observation path is not a regular file")
}
if wait == nil {
return 0, errors.New("observation quiet-window wait is nil")
}
if err := wait(); err != nil {
return 0, err
}
after, err := os.Stat(path)
if err != nil {
return 0, err
}
if !after.Mode().IsRegular() {
return 0, errors.New("observation path stopped being a regular file")
}
if !os.SameFile(before, after) {
return 0, errors.New("observation file was rotated during the quiet window")
}
if after.Size() != before.Size() {
return 0, fmt.Errorf(
"observation file size changed during the quiet window: before=%d after=%d",
before.Size(), after.Size(),
)
}
return after.Size(), nil
}
func waitDevRepeatGuardObservationSegment(
path string,
offset int64,
model string,
expectedCorrelations int,
forbiddenValues ...[]byte,
) ([]devRepeatGuardObservation, error) {
deadline := time.Now().Add(15 * time.Second)
var lastErr error
for {
observations, err := readDevRepeatGuardObservationSegment(path, offset, forbiddenValues...)
if err == nil {
groups := groupDevRepeatGuardObservations(observations, model)
if len(groups) > expectedCorrelations {
return nil, fmt.Errorf(
"observation segment is contaminated: groups=%d, want exactly %d",
len(groups), expectedCorrelations,
)
}
complete := 0
for _, group := range groups {
for _, observation := range group {
if observation.ObservationKind == string(streamgate.ObservationKindTerminalCommitted) {
complete++
break
}
}
}
if len(groups) == expectedCorrelations && complete == expectedCorrelations {
return observations, nil
}
lastErr = fmt.Errorf(
"observation groups=%d complete=%d, want %d",
len(groups), complete, expectedCorrelations,
)
} else {
lastErr = err
}
if time.Now().After(deadline) {
return nil, lastErr
}
time.Sleep(250 * time.Millisecond)
}
}
func readDevRepeatGuardObservationSegment(
path string,
offset int64,
forbiddenValues ...[]byte,
) ([]devRepeatGuardObservation, error) {
file, err := os.Open(path)
if err != nil {
return nil, err
}
defer file.Close()
info, err := file.Stat()
if err != nil {
return nil, err
}
if info.Size() < offset {
return nil, errors.New("observation file was truncated or rotated")
}
if _, err := file.Seek(offset, io.SeekStart); err != nil {
return nil, err
}
const observationSegmentLimit = 64 * 1024 * 1024
raw, err := io.ReadAll(io.LimitReader(file, observationSegmentLimit+1))
if err != nil {
return nil, err
}
if len(raw) > observationSegmentLimit {
return nil, errors.New("observation segment exceeds 64 MiB")
}
return parseDevRepeatGuardObservations(raw, forbiddenValues...)
}
func parseDevRepeatGuardObservations(
raw []byte,
forbiddenValues ...[]byte,
) ([]devRepeatGuardObservation, error) {
for _, forbidden := range forbiddenValues {
if len(forbidden) > 0 && bytes.Contains(raw, forbidden) {
return nil, errors.New("observation segment contains a forbidden raw value")
}
}
observations := make([]devRepeatGuardObservation, 0)
for lineIndex, rawLine := range bytes.Split(raw, []byte{'\n'}) {
line := bytes.TrimSpace(rawLine)
if len(line) == 0 {
continue
}
var fields map[string]json.RawMessage
if err := json.Unmarshal(line, &fields); err != nil {
return nil, fmt.Errorf("observation line %d is not JSON", lineIndex+1)
}
if err := validateDevRepeatGuardObservationFields(fields); err != nil {
return nil, fmt.Errorf("observation line %d: %w", lineIndex+1, err)
}
var observation devRepeatGuardObservation
if err := json.Unmarshal(line, &observation); err != nil {
return nil, fmt.Errorf("decode observation line %d: %w", lineIndex+1, err)
}
if observation.Message != filterObservationLogMessage {
continue
}
observation.line = lineIndex
if strings.TrimSpace(observation.CorrelationID) == "" {
return nil, fmt.Errorf("observation line %d has no correlation_id", lineIndex+1)
}
observations = append(observations, observation)
}
if len(observations) == 0 {
return nil, errors.New("observation segment has no streamgate filter observations")
}
return observations, nil
}
func validateDevRepeatGuardObservationFields(fields map[string]json.RawMessage) error {
forbidden := map[string]struct{}{
"authorization": {}, "api_key": {}, "token": {}, "prompt": {},
"raw_prompt": {}, "output": {}, "raw_output": {}, "arguments": {},
"tool_arguments": {}, "tool_result": {}, "request_body": {}, "response_body": {},
}
for key := range fields {
if _, found := forbidden[strings.ToLower(strings.TrimSpace(key))]; found {
return fmt.Errorf("forbidden raw field %q", key)
}
}
return nil
}
func groupDevRepeatGuardObservations(
observations []devRepeatGuardObservation,
model string,
) map[string][]devRepeatGuardObservation {
groups := make(map[string][]devRepeatGuardObservation)
for _, observation := range observations {
if observation.ModelGroup != model {
continue
}
groups[observation.CorrelationID] = append(groups[observation.CorrelationID], observation)
}
return groups
}
func correlateDevRepeatGuardBatch(
batch []devRepeatGuardSmokeEvidence,
observations []devRepeatGuardObservation,
) ([]devRepeatGuardLifecycleEvidence, error) {
if len(batch) == 0 {
return nil, errors.New("no request results to correlate")
}
run := batch[0].Run
model := batch[0].Model
providerResponseIDs := make(map[string]struct{}, len(batch))
attempts := make(map[int]struct{}, len(batch))
providerTerminalCounts := map[string]int{"done": 0, "error": 0}
for _, response := range batch {
if response.Run != run || response.Model != model {
return nil, errors.New("provider response batch mixes run or model identity")
}
if response.Attempt <= 0 {
return nil, errors.New("provider response has no attempt index")
}
if _, found := attempts[response.Attempt]; found {
return nil, fmt.Errorf("duplicate provider response attempt=%d", response.Attempt)
}
attempts[response.Attempt] = struct{}{}
if response.StatusCode < http.StatusOK || response.StatusCode >= http.StatusMultipleChoices {
return nil, fmt.Errorf(
"run=%d attempt=%d returned non-2xx status %d",
response.Run, response.Attempt, response.StatusCode,
)
}
if _, found := providerTerminalCounts[response.Terminal]; !found {
return nil, fmt.Errorf(
"run=%d attempt=%d has invalid SSE terminal %q",
response.Run, response.Attempt, response.Terminal,
)
}
providerTerminalCounts[response.Terminal]++
if err := collectDevRepeatGuardProviderResponseIDs(
response.ProviderResponseIDs,
providerResponseIDs,
fmt.Sprintf("run=%d attempt=%d", response.Run, response.Attempt),
); err != nil {
return nil, err
}
if response.downstreamRepeated {
return nil, fmt.Errorf(
"run=%d attempt=%d released repeated downstream content",
response.Run, response.Attempt,
)
}
}
groups := groupDevRepeatGuardObservations(observations, model)
if len(groups) != len(batch) {
return nil, fmt.Errorf("correlation groups=%d, requests=%d", len(groups), len(batch))
}
correlations := make([]string, 0, len(groups))
for correlation := range groups {
correlations = append(correlations, correlation)
}
sort.Strings(correlations)
edgeTerminalCounts := map[string]int{"done": 0, "error": 0}
correlated := make([]devRepeatGuardLifecycleEvidence, 0, len(groups))
for _, correlation := range correlations {
evidence, err := correlateRepeatGuardObservations(run, model, groups[correlation])
if err != nil {
return nil, fmt.Errorf(
"run=%d correlation=%s: %w",
run, correlation, err,
)
}
terminal := "done"
if evidence.TerminalReason == "error" {
terminal = "error"
}
edgeTerminalCounts[terminal]++
correlated = append(correlated, evidence)
}
if providerTerminalCounts["done"] != edgeTerminalCounts["done"] ||
providerTerminalCounts["error"] != edgeTerminalCounts["error"] {
return nil, fmt.Errorf(
"provider/Edge terminal cardinality mismatch: provider done/error=%d/%d Edge done/error=%d/%d",
providerTerminalCounts["done"], providerTerminalCounts["error"],
edgeTerminalCounts["done"], edgeTerminalCounts["error"],
)
}
return correlated, nil
}
func correlateRepeatGuardObservations(
run int,
model string,
observations []devRepeatGuardObservation,
) (devRepeatGuardLifecycleEvidence, error) {
if len(observations) == 0 {
return devRepeatGuardLifecycleEvidence{}, errors.New("no correlated observations")
}
correlation := observations[0].CorrelationID
if strings.TrimSpace(correlation) == "" {
return devRepeatGuardLifecycleEvidence{}, errors.New("missing correlation_id")
}
var (
actualModel, actualProvider string
repeatPass, repeatViolation bool
repeatFatal, arbitrationRecover bool
repeatEvidence bool
arbitrationSeen int
arbitrationTerminal, terminalSeen int
terminalReason string
recoveryKinds = make(map[string]int)
recoveryStage int
recoveryLifecycleComplete bool
fingerprint string
offset int
)
for _, observation := range observations {
if observation.CorrelationID != correlation {
return devRepeatGuardLifecycleEvidence{}, errors.New("mixed correlation IDs")
}
if observation.ActualModel != "" {
actualModel = observation.ActualModel
}
if observation.ActualProvider != "" {
actualProvider = observation.ActualProvider
}
switch observation.ObservationKind {
case string(streamgate.ObservationKindFilterEvaluated):
if observation.FilterID != openAIRepeatGuardFilterID &&
observation.FilterID != openAIRepeatActionGuardFilterID {
continue
}
if observation.FilterOutcome != string(streamgate.FilterOutcomeKindEvaluated) {
continue
}
switch observation.DecisionKind {
case string(streamgate.FilterDecisionKindPass):
repeatPass = true
case string(streamgate.FilterDecisionKindViolation):
repeatViolation = true
case string(streamgate.FilterDecisionKindFatal):
repeatFatal = true
}
if observation.EvidenceFingerprint != "" &&
observation.DecisionKind != string(streamgate.FilterDecisionKindPass) {
fingerprint = observation.EvidenceFingerprint
offset = observation.EvidenceOffset
}
switch observation.EvidenceDescriptorCode {
case "repeat_content_detected",
"assistant_history_anchor_repeat",
"repeated_action_no_progress",
"repeat_after_tool_boundary",
"repeat_recovery_source_unavailable":
repeatEvidence = true
}
case string(streamgate.ObservationKindArbitrationDecided):
arbitrationSeen++
switch observation.ArbitrationAction {
case string(streamgate.ArbitrationActionRecover):
arbitrationRecover = true
case string(streamgate.ArbitrationActionTerminal):
arbitrationTerminal++
}
case string(streamgate.ObservationKindTerminalCommitted):
terminalSeen++
switch observation.TerminalReason {
case "":
terminalReason = "error"
case string(streamgate.TerminalReasonCompleted):
terminalReason = observation.TerminalReason
default:
return devRepeatGuardLifecycleEvidence{}, fmt.Errorf(
"unexpected terminal reason %q", observation.TerminalReason,
)
}
case string(streamgate.ObservationKindRecoveryPlanSelected),
string(streamgate.ObservationKindRecoveryAttemptAborted),
string(streamgate.ObservationKindRecoveryRebuilt),
string(streamgate.ObservationKindRecoveryDispatched):
recoveryKinds[observation.ObservationKind]++
switch observation.ObservationKind {
case string(streamgate.ObservationKindRecoveryPlanSelected):
recoveryStage = 1
case string(streamgate.ObservationKindRecoveryAttemptAborted):
if recoveryStage == 1 {
recoveryStage = 2
}
case string(streamgate.ObservationKindRecoveryRebuilt):
if recoveryStage == 2 {
recoveryStage = 3
}
case string(streamgate.ObservationKindRecoveryDispatched):
if recoveryStage == 3 {
recoveryLifecycleComplete = true
}
}
if observation.RecoveryStrategy != "" &&
observation.RecoveryStrategy != string(streamgate.RecoveryStrategyContinuationRepair) {
return devRepeatGuardLifecycleEvidence{}, fmt.Errorf(
"unexpected recovery strategy %q", observation.RecoveryStrategy,
)
}
}
}
if actualModel == "" {
return devRepeatGuardLifecycleEvidence{}, errors.New("missing actual_model")
}
if actualProvider == "" || actualProvider == "unknown" || actualProvider == "unavailable" {
return devRepeatGuardLifecycleEvidence{}, errors.New("missing actual_provider")
}
if arbitrationSeen == 0 {
return devRepeatGuardLifecycleEvidence{}, errors.New("missing arbitration_decided")
}
if terminalSeen != 1 {
return devRepeatGuardLifecycleEvidence{}, fmt.Errorf(
"terminal_committed count=%d, want 1", terminalSeen,
)
}
evidence := devRepeatGuardLifecycleEvidence{
Run: run, Model: model, ActualModel: actualModel, Provider: actualProvider,
CorrelationID: correlation, TerminalReason: terminalReason,
Fingerprint: fingerprint, Offset: offset,
}
switch {
case repeatFatal:
if !repeatEvidence {
return devRepeatGuardLifecycleEvidence{}, errors.New("fatal guard decision has no repeat evidence")
}
if arbitrationTerminal == 0 {
return devRepeatGuardLifecycleEvidence{}, errors.New("fatal repeat has no terminal arbitration")
}
if terminalReason != "error" {
return devRepeatGuardLifecycleEvidence{}, errors.New("fatal repeat did not end with an Edge error terminal")
}
evidence.Result = "reproduced"
evidence.Decision = "safe_stop"
evidence.Lifecycle = "evaluated_fatal_terminal_committed"
case repeatViolation:
if !repeatEvidence {
return devRepeatGuardLifecycleEvidence{}, errors.New("guard violation has no repeat evidence")
}
if !arbitrationRecover {
return devRepeatGuardLifecycleEvidence{}, errors.New("repeat violation has no recover arbitration")
}
if terminalReason == "error" {
for _, kind := range []streamgate.ObservationKind{
streamgate.ObservationKindRecoveryPlanSelected,
streamgate.ObservationKindRecoveryAttemptAborted,
} {
if recoveryKinds[string(kind)] == 0 {
return devRepeatGuardLifecycleEvidence{}, fmt.Errorf("repeat safe stop missing %s", kind)
}
}
evidence.Result = "reproduced"
evidence.Decision = "safe_stop"
evidence.Lifecycle = "evaluated_violation_abort_terminal_committed"
break
}
if terminalReason != string(streamgate.TerminalReasonCompleted) {
return devRepeatGuardLifecycleEvidence{}, errors.New("repeat recovery has no completed Edge terminal")
}
required := []streamgate.ObservationKind{
streamgate.ObservationKindRecoveryPlanSelected,
streamgate.ObservationKindRecoveryAttemptAborted,
streamgate.ObservationKindRecoveryRebuilt,
streamgate.ObservationKindRecoveryDispatched,
}
for _, kind := range required {
if recoveryKinds[string(kind)] == 0 {
return devRepeatGuardLifecycleEvidence{}, fmt.Errorf("repeat violation missing %s", kind)
}
}
if !recoveryLifecycleComplete {
return devRepeatGuardLifecycleEvidence{}, errors.New("repeat recovery lifecycle is out of order")
}
evidence.Result = "reproduced"
evidence.Decision = "continuation_repair"
evidence.Lifecycle = "evaluated_violation_abort_rebuild_dispatch_terminal_committed"
case repeatPass:
if terminalReason != string(streamgate.TerminalReasonCompleted) {
return devRepeatGuardLifecycleEvidence{}, errors.New("passing repeat guard did not end successfully")
}
evidence.Result = "not_reproduced"
evidence.Decision = "guard_evaluated_pass"
evidence.Lifecycle = "evaluated_pass_terminal_committed"
default:
return devRepeatGuardLifecycleEvidence{}, errors.New("repeat guard was not evaluated")
}
return evidence, nil
}
func TestDevRepeatGuardSmokeEvidence(t *testing.T) {
const (
errorEvent = "data: {\"error\":{\"type\":\"run_error\"}}\n\n"
doneEvent = "data: [DONE]\n\n"
)
providerEvent := func(id, content string) string {
return fmt.Sprintf(
"data: {\"id\":%q,\"choices\":[{\"delta\":{\"content\":%q},\"finish_reason\":null}]}\n\n",
id, content,
)
}
successFinish := func(id string) string {
return fmt.Sprintf(
"data: {\"id\":%q,\"choices\":[{\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n",
id,
)
}
validSuccess := providerEvent("chatcmpl-run-1", "ok") +
successFinish("chatcmpl-run-1") + doneEvent
validContinuationSuccess := providerEvent("chatcmpl-run-1", "first") +
providerEvent("chatcmpl-run-2", "continued") +
successFinish("chatcmpl-run-2") + doneEvent
validContinuationError := providerEvent("chatcmpl-run-1", "first") +
providerEvent("chatcmpl-run-2", "continued") + errorEvent + doneEvent
tests := []struct {
name string
status int
body string
wantTerminal string
wantIDs []string
wantErr bool
}{
{
name: "valid one identity success", status: http.StatusOK,
body: validSuccess, wantTerminal: "done",
wantIDs: []string{"chatcmpl-run-1"},
},
{
name: "valid two identity continuation success", status: http.StatusOK,
body: validContinuationSuccess, wantTerminal: "done",
wantIDs: []string{"chatcmpl-run-1", "chatcmpl-run-2"},
},
{
name: "valid two identity continuation error", status: http.StatusOK,
body: validContinuationError, wantTerminal: "error",
wantIDs: []string{"chatcmpl-run-1", "chatcmpl-run-2"},
},
{name: "non-2xx", status: http.StatusBadGateway, body: validSuccess, wantErr: true},
{name: "malformed JSON", status: http.StatusOK, body: "data: {\n\ndata: [DONE]\n\n", wantErr: true},
{
name: "missing provider response identity", status: http.StatusOK,
body: "data: {\"choices\":[{\"finish_reason\":\"stop\"}]}\n\n" + doneEvent, wantErr: true,
},
{
name: "empty provider response identity", status: http.StatusOK,
body: "data: {\"id\":\"\",\"choices\":[{\"finish_reason\":\"stop\"}]}\n\n" + doneEvent,
wantErr: true,
},
{
name: "invalid provider response identity prefix", status: http.StatusOK,
body: "data: {\"id\":\"resp-run-1\",\"choices\":[{\"finish_reason\":\"stop\"}]}\n\n" + doneEvent,
wantErr: true,
},
{
name: "non contiguous provider response identity reuse", status: http.StatusOK,
body: providerEvent("chatcmpl-run-1", "first") +
providerEvent("chatcmpl-run-2", "continued") +
successFinish("chatcmpl-run-1") + doneEvent,
wantErr: true,
},
{
name: "more than four provider response identities", status: http.StatusOK,
body: providerEvent("chatcmpl-run-1", "one") +
providerEvent("chatcmpl-run-2", "two") +
providerEvent("chatcmpl-run-3", "three") +
providerEvent("chatcmpl-run-4", "four") +
successFinish("chatcmpl-run-5") + doneEvent,
wantErr: true,
},
{
name: "missing success finish", status: http.StatusOK,
body: providerEvent("chatcmpl-run-1", "ok") + doneEvent, wantErr: true,
},
{
name: "missing success done", status: http.StatusOK,
body: providerEvent("chatcmpl-run-1", "ok") + successFinish("chatcmpl-run-1"),
wantErr: true,
},
{
name: "missing error done", status: http.StatusOK,
body: providerEvent("chatcmpl-run-1", "ok") + errorEvent, wantErr: true,
},
{
name: "duplicate error", status: http.StatusOK,
body: providerEvent("chatcmpl-run-1", "ok") + errorEvent + errorEvent + doneEvent,
wantErr: true,
},
{
name: "duplicate success finish", status: http.StatusOK,
body: providerEvent("chatcmpl-run-1", "ok") +
successFinish("chatcmpl-run-1") + successFinish("chatcmpl-run-1") + doneEvent,
wantErr: true,
},
{name: "duplicate done", status: http.StatusOK, body: validSuccess + doneEvent, wantErr: true},
{
name: "event after done", status: http.StatusOK,
body: validSuccess + providerEvent("chatcmpl-run-1", "late"), wantErr: true,
},
{
name: "data after logical terminal", status: http.StatusOK,
body: providerEvent("chatcmpl-run-1", "ok") +
successFinish("chatcmpl-run-1") +
providerEvent("chatcmpl-run-1", "late") + doneEvent,
wantErr: true,
},
{
name: "finish then error", status: http.StatusOK,
body: providerEvent("chatcmpl-run-1", "ok") +
successFinish("chatcmpl-run-1") + errorEvent + doneEvent,
wantErr: true,
},
{
name: "error then finish", status: http.StatusOK,
body: providerEvent("chatcmpl-run-1", "ok") +
errorEvent + successFinish("chatcmpl-run-1") + doneEvent,
wantErr: true,
},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
w.WriteHeader(tc.status)
_, _ = io.WriteString(w, tc.body)
}))
defer server.Close()
artifactDir := t.TempDir()
evidence, err := runDevRepeatGuardSmokeAttempt(
server.Client(), server.URL, "ornith:35b", "test prompt",
artifactDir, 1, 1,
)
if (err != nil) != tc.wantErr {
t.Fatalf("runDevRepeatGuardSmokeAttempt error=%v, wantErr=%v", err, tc.wantErr)
}
if !tc.wantErr {
if evidence.StatusCode != http.StatusOK || evidence.Terminal != tc.wantTerminal {
t.Fatalf("evidence=%+v", evidence)
}
if !reflect.DeepEqual(evidence.ProviderResponseIDs, tc.wantIDs) {
t.Fatalf("provider response identities=%v, want %v", evidence.ProviderResponseIDs, tc.wantIDs)
}
}
})
}
}
func TestDevRepeatGuardDownstreamRepeatEvidence(t *testing.T) {
const minimum = 64
contentToken := strings.Repeat("C", minimum)
reasoningToken := strings.Repeat("R", minimum)
aliasToken := strings.Repeat("A", minimum)
if len([]rune(contentToken)) != minimum ||
len([]rune(reasoningToken)) != minimum ||
len([]rune(aliasToken)) != minimum {
t.Fatalf("fixture tokens must each be %d runes", minimum)
}
deltaEvent := func(fields string) string {
return "data: {\"choices\":[{\"delta\":{" + fields + "}}]}\n\n"
}
contentField := func(value string) string {
return fmt.Sprintf("%q:%q", "content", value)
}
reasoningContentField := func(value string) string {
return fmt.Sprintf("%q:%q", "reasoning_content", value)
}
reasoningField := func(value string) string {
return fmt.Sprintf("%q:%q", "reasoning", value)
}
const done = "data: [DONE]\n\n"
t.Run("cross-channel identical tails are not a repeat", func(t *testing.T) {
raw := deltaEvent(contentField(contentToken)) +
deltaEvent(reasoningContentField(contentToken)) + done
channels := devRepeatGuardSmokeChannels([]byte(raw))
if channels.contentRepeated(minimum) || channels.reasoningRepeated(minimum) {
t.Fatalf("identical tails in distinct channels must not be a downstream repeat")
}
})
t.Run("same-channel content repeat is detected", func(t *testing.T) {
raw := deltaEvent(contentField(contentToken)) +
deltaEvent(contentField(contentToken)) + done
channels := devRepeatGuardSmokeChannels([]byte(raw))
if !channels.contentRepeated(minimum) {
t.Fatalf("duplicated content channel must be a downstream repeat")
}
if channels.reasoningRepeated(minimum) {
t.Fatalf("empty reasoning channel must not repeat")
}
})
t.Run("same-channel reasoning repeat is detected", func(t *testing.T) {
raw := deltaEvent(reasoningContentField(reasoningToken)) +
deltaEvent(reasoningContentField(reasoningToken)) + done
channels := devRepeatGuardSmokeChannels([]byte(raw))
if !channels.reasoningRepeated(minimum) {
t.Fatalf("duplicated reasoning channel must be a downstream repeat")
}
if channels.contentRepeated(minimum) {
t.Fatalf("empty content channel must not repeat")
}
})
t.Run("reasoning alias precedence uses reasoning_content only", func(t *testing.T) {
raw := deltaEvent(reasoningContentField(reasoningToken)+","+reasoningField(aliasToken)) +
deltaEvent(reasoningContentField(reasoningToken)+","+reasoningField(aliasToken)) + done
channels := devRepeatGuardSmokeChannels([]byte(raw))
if got, want := string(channels.reasoning), reasoningToken+reasoningToken; got != want {
t.Fatalf("reasoning channel must use reasoning_content precedence only")
}
if strings.Contains(string(channels.reasoning), aliasToken) {
t.Fatalf("reasoning alias must not be appended alongside reasoning_content")
}
if !channels.reasoningRepeated(minimum) {
t.Fatalf("duplicated reasoning_content must be a downstream repeat")
}
})
t.Run("single event does not concatenate reasoning aliases", func(t *testing.T) {
raw := deltaEvent(reasoningContentField(reasoningToken)+","+reasoningField(reasoningToken)) + done
channels := devRepeatGuardSmokeChannels([]byte(raw))
if channels.reasoningRepeated(minimum) {
t.Fatalf("a single reasoning value must not be doubled by alias concatenation")
}
})
}
func TestDevRepeatGuardRequestDeadline(t *testing.T) {
now := time.Unix(1_800_000_000, 0)
exactThreeRunBudget := 3*
(devRepeatGuardDefaultRequestTimeout+devRepeatGuardPerRunSettlementAllowance) +
devRepeatGuardFinalCleanupReserve
tests := []struct {
name string
raw string
runs int
deadline time.Time
hasDeadline bool
want time.Duration
wantErr bool
}{
{
name: "default duration without test deadline", runs: 3,
want: devRepeatGuardDefaultRequestTimeout,
},
{
name: "explicit valid duration", raw: "24m", runs: 3,
want: 24 * time.Minute,
},
{name: "malformed duration", raw: "twenty-five", runs: 3, wantErr: true},
{name: "zero duration", raw: "0s", runs: 3, wantErr: true},
{name: "negative duration", raw: "-1s", runs: 3, wantErr: true},
{name: "above policy duration", raw: "25m1s", runs: 3, wantErr: true},
{
name: "exact fit includes settlement and cleanup", raw: "25m", runs: 3,
deadline: now.Add(exactThreeRunBudget), hasDeadline: true,
want: devRepeatGuardDefaultRequestTimeout,
},
{
name: "insufficient final cleanup reserve", raw: "25m", runs: 3,
deadline: now.Add(exactThreeRunBudget - time.Nanosecond), hasDeadline: true,
wantErr: true,
},
{
name: "missing deadline falls back to policy validation", raw: "25m", runs: 3,
deadline: now.Add(time.Second), hasDeadline: false,
want: devRepeatGuardDefaultRequestTimeout,
},
{
name: "three runs fit a ninety minute envelope", raw: "25m", runs: 3,
deadline: now.Add(90 * time.Minute), hasDeadline: true,
want: devRepeatGuardDefaultRequestTimeout,
},
{name: "invalid run count", raw: "25m", runs: 0, wantErr: true},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
got, err := resolveDevRepeatGuardRequestTimeout(
tc.raw, tc.runs, now, tc.deadline, tc.hasDeadline,
)
if (err != nil) != tc.wantErr {
t.Fatalf("resolve error=%v, wantErr=%v", err, tc.wantErr)
}
if !tc.wantErr && got != tc.want {
t.Fatalf("request timeout=%s, want %s", got, tc.want)
}
})
}
}
func TestDevRepeatGuardObservationCorrelation(t *testing.T) {
base := devRepeatGuardSmokeEvidence{
Run: 1, Attempt: 1, Model: "ornith:35b", StatusCode: http.StatusOK,
Terminal: "done", ProviderResponseIDs: []string{"chatcmpl-provider-1"},
}
observation := func(kind string) devRepeatGuardObservation {
return devRepeatGuardObservation{
Message: filterObservationLogMessage, ObservationKind: kind,
CorrelationID: "req.run-1", ModelGroup: "ornith:35b",
ActualModel: "served-model", ActualProvider: "provider-a",
}
}
filterDecision := func(decision string) devRepeatGuardObservation {
value := observation(string(streamgate.ObservationKindFilterEvaluated))
value.FilterID = openAIRepeatGuardFilterID
value.FilterOutcome = string(streamgate.FilterOutcomeKindEvaluated)
value.DecisionKind = decision
value.EvidenceFingerprint = "0123456789abcdef"
value.EvidenceOffset = 23
if decision != string(streamgate.FilterDecisionKindPass) {
value.EvidenceDescriptorCode = "repeat_content_detected"
}
return value
}
terminal := observation(string(streamgate.ObservationKindTerminalCommitted))
terminal.TerminalReason = string(streamgate.TerminalReasonCompleted)
errorTerminal := observation(string(streamgate.ObservationKindTerminalCommitted))
t.Run("pass is not reproduced", func(t *testing.T) {
evidence, err := correlateRepeatGuardObservations(1, base.Model, []devRepeatGuardObservation{
filterDecision(string(streamgate.FilterDecisionKindPass)),
observation(string(streamgate.ObservationKindArbitrationDecided)),
terminal,
})
if err != nil {
t.Fatal(err)
}
if evidence.Result != "not_reproduced" || evidence.Decision != "guard_evaluated_pass" {
t.Fatalf("evidence=%+v", evidence)
}
})
t.Run("violation requires continuation lifecycle", func(t *testing.T) {
violation := filterDecision(string(streamgate.FilterDecisionKindViolation))
arbitration := observation(string(streamgate.ObservationKindArbitrationDecided))
arbitration.ArbitrationAction = string(streamgate.ArbitrationActionRecover)
lifecycle := []devRepeatGuardObservation{violation, arbitration}
for _, kind := range []streamgate.ObservationKind{
streamgate.ObservationKindRecoveryPlanSelected,
streamgate.ObservationKindRecoveryAttemptAborted,
streamgate.ObservationKindRecoveryRebuilt,
streamgate.ObservationKindRecoveryDispatched,
} {
value := observation(string(kind))
value.RecoveryStrategy = string(streamgate.RecoveryStrategyContinuationRepair)
lifecycle = append(lifecycle, value)
}
lifecycle = append(lifecycle, terminal)
evidence, err := correlateRepeatGuardObservations(1, base.Model, lifecycle)
if err != nil {
t.Fatal(err)
}
if evidence.Result != "reproduced" || evidence.Decision != "continuation_repair" {
t.Fatalf("evidence=%+v", evidence)
}
})
t.Run("violation can end in safe stop", func(t *testing.T) {
violation := filterDecision(string(streamgate.FilterDecisionKindViolation))
arbitration := observation(string(streamgate.ObservationKindArbitrationDecided))
arbitration.ArbitrationAction = string(streamgate.ArbitrationActionRecover)
selected := observation(string(streamgate.ObservationKindRecoveryPlanSelected))
selected.RecoveryStrategy = string(streamgate.RecoveryStrategyContinuationRepair)
aborted := observation(string(streamgate.ObservationKindRecoveryAttemptAborted))
aborted.RecoveryStrategy = string(streamgate.RecoveryStrategyContinuationRepair)
evidence, err := correlateRepeatGuardObservations(1, base.Model, []devRepeatGuardObservation{
violation, arbitration, selected, aborted, errorTerminal,
})
if err != nil {
t.Fatal(err)
}
if evidence.Result != "reproduced" || evidence.Decision != "safe_stop" {
t.Fatalf("evidence=%+v", evidence)
}
})
t.Run("fatal is safe stop", func(t *testing.T) {
arbitration := observation(string(streamgate.ObservationKindArbitrationDecided))
arbitration.ArbitrationAction = string(streamgate.ArbitrationActionTerminal)
evidence, err := correlateRepeatGuardObservations(1, base.Model, []devRepeatGuardObservation{
filterDecision(string(streamgate.FilterDecisionKindFatal)), arbitration, errorTerminal,
})
if err != nil {
t.Fatal(err)
}
if evidence.Result != "reproduced" || evidence.Decision != "safe_stop" {
t.Fatalf("evidence=%+v", evidence)
}
})
for _, tc := range []struct {
name string
mutate func([]devRepeatGuardObservation) []devRepeatGuardObservation
}{
{
name: "missing provider",
mutate: func(values []devRepeatGuardObservation) []devRepeatGuardObservation {
for i := range values {
values[i].ActualProvider = ""
}
return values
},
},
{
name: "missing correlation",
mutate: func(values []devRepeatGuardObservation) []devRepeatGuardObservation {
for i := range values {
values[i].CorrelationID = ""
}
return values
},
},
} {
t.Run(tc.name, func(t *testing.T) {
values := tc.mutate([]devRepeatGuardObservation{
filterDecision(string(streamgate.FilterDecisionKindPass)), terminal,
})
if _, err := correlateRepeatGuardObservations(1, base.Model, values); err == nil {
t.Fatal("expected correlation failure")
}
})
}
t.Run("raw sentinel is rejected", func(t *testing.T) {
raw := []byte(
"{\"msg\":\"streamgate_filter_observation\",\"correlation_id\":\"req.run-1\"," +
"\"model_group\":\"ornith:35b\",\"prompt\":\"raw-sentinel\"}\n",
)
if _, err := parseDevRepeatGuardObservations(raw, []byte("raw-sentinel")); err == nil {
t.Fatal("expected raw sentinel rejection")
}
})
passGroup := func(correlation string) []devRepeatGuardObservation {
filter := filterDecision(string(streamgate.FilterDecisionKindPass))
arbitration := observation(string(streamgate.ObservationKindArbitrationDecided))
committed := terminal
for _, value := range []*devRepeatGuardObservation{&filter, &arbitration, &committed} {
value.CorrelationID = correlation
}
return []devRepeatGuardObservation{filter, arbitration, committed}
}
request := func(attempt int, providerResponseIDs ...string) devRepeatGuardSmokeEvidence {
value := base
value.Attempt = attempt
value.ProviderResponseIDs = append([]string(nil), providerResponseIDs...)
return value
}
t.Run("batch accepts disjoint provider and Edge identities", func(t *testing.T) {
batch := []devRepeatGuardSmokeEvidence{
request(1, "chatcmpl-provider-a", "chatcmpl-provider-a-continuation"),
request(2, "chatcmpl-provider-b"),
}
observations := append(passGroup("req.edge-z"), passGroup("req.edge-a")...)
correlated, err := correlateDevRepeatGuardBatch(batch, observations)
if err != nil {
t.Fatal(err)
}
if correlated[0].CorrelationID != "req.edge-a" ||
correlated[1].CorrelationID != "req.edge-z" {
t.Fatalf("correlated=%+v", correlated)
}
})
t.Run("final evidence rejects overlap across runs", func(t *testing.T) {
firstRun := request(1, "chatcmpl-provider-a", "chatcmpl-provider-shared")
secondRun := request(1, "chatcmpl-provider-shared", "chatcmpl-provider-b")
secondRun.Run = 2
err := validateDevRepeatGuardUniqueEvidence(
[]devRepeatGuardSmokeEvidence{firstRun, secondRun},
[]devRepeatGuardLifecycleEvidence{
{CorrelationID: "req.edge-a"},
{CorrelationID: "req.edge-b"},
},
)
if err == nil {
t.Fatal("expected final flattened provider identity overlap to fail")
}
})
for _, tc := range []struct {
name string
batch []devRepeatGuardSmokeEvidence
observations []devRepeatGuardObservation
}{
{
name: "batch rejects missing provider identity",
batch: []devRepeatGuardSmokeEvidence{request(1)},
observations: passGroup("req.run-1"),
},
{
name: "batch rejects duplicate provider identity",
batch: []devRepeatGuardSmokeEvidence{
request(1, "chatcmpl-provider-a"),
request(2, "chatcmpl-provider-a"),
},
observations: append(passGroup("req.run-1"), passGroup("req.run-2")...),
},
{
name: "batch rejects cross-request provider identity overlap",
batch: []devRepeatGuardSmokeEvidence{
request(1, "chatcmpl-provider-a", "chatcmpl-provider-shared"),
request(2, "chatcmpl-provider-shared", "chatcmpl-provider-b"),
},
observations: append(passGroup("req.run-1"), passGroup("req.run-2")...),
},
{
name: "batch rejects extra Edge group",
batch: []devRepeatGuardSmokeEvidence{
request(1, "chatcmpl-provider-a"),
request(2, "chatcmpl-provider-b"),
},
observations: append(
append(passGroup("req.run-1"), passGroup("req.run-2")...),
passGroup("req.manual")...,
),
},
{
name: "batch rejects provider and Edge terminal count mismatch",
batch: []devRepeatGuardSmokeEvidence{
request(1, "chatcmpl-provider-a"),
func() devRepeatGuardSmokeEvidence {
value := request(2, "chatcmpl-provider-b")
value.Terminal = "error"
return value
}(),
},
observations: append(passGroup("req.run-1"), passGroup("req.run-2")...),
},
{
name: "batch rejects released repeat without identity pairing",
batch: []devRepeatGuardSmokeEvidence{
func() devRepeatGuardSmokeEvidence {
value := request(1, "chatcmpl-provider-a")
value.downstreamRepeated = true
return value
}(),
},
observations: passGroup("req.run-1"),
},
} {
t.Run(tc.name, func(t *testing.T) {
if _, err := correlateDevRepeatGuardBatch(tc.batch, tc.observations); err == nil {
t.Fatal("expected strict batch correlation failure")
}
})
}
}
func TestDevRepeatGuardObservationQuiescence(t *testing.T) {
path := filepath.Join(t.TempDir(), "observations.jsonl")
const initial = "{\"msg\":\"sanitized\"}\n"
if err := os.WriteFile(path, []byte(initial), 0o600); err != nil {
t.Fatal(err)
}
t.Run("unchanged regular file passes", func(t *testing.T) {
offset, err := devRepeatGuardObservationQuietOffsetWithWait(path, func() error {
return nil
})
if err != nil {
t.Fatal(err)
}
if offset != int64(len(initial)) {
t.Fatalf("offset=%d, want %d", offset, len(initial))
}
})
t.Run("append during quiet window fails", func(t *testing.T) {
if _, err := devRepeatGuardObservationQuietOffsetWithWait(path, func() error {
file, openErr := os.OpenFile(path, os.O_APPEND|os.O_WRONLY, 0)
if openErr != nil {
return openErr
}
if _, writeErr := file.WriteString("{\"msg\":\"concurrent\"}\n"); writeErr != nil {
_ = file.Close()
return writeErr
}
return file.Close()
}); err == nil {
t.Fatal("expected append during quiet window to fail")
}
})
t.Run("rotation during quiet window fails", func(t *testing.T) {
if _, err := devRepeatGuardObservationQuietOffsetWithWait(path, func() error {
rotated := path + ".rotated"
if renameErr := os.Rename(path, rotated); renameErr != nil {
return renameErr
}
return os.WriteFile(path, []byte(initial), 0o600)
}); err == nil {
t.Fatal("expected rotation during quiet window to fail")
}
})
}
func TestDevRepeatGuardCapacityEvidence(t *testing.T) {
const validStatus = `{
"nodes": [
{"provider_snapshots": [
{"id":"provider-a","capacity":2,"in_flight":2,"queued":1},
{"id":"ignored-provider","capacity":99,"in_flight":99,"queued":99}
]},
{"provider_snapshots": [
{"id":"provider-b","capacity":2,"in_flight":2,"queued":2}
]}
]
}`
providerIDs, err := parseDevRepeatGuardProviderIDs(" provider-a,provider-b ")
if err != nil {
t.Fatal(err)
}
snapshot, err := parseDevRepeatGuardCapacityStatus([]byte(validStatus), providerIDs)
if err != nil {
t.Fatal(err)
}
if snapshot.configuredCapacity != 4 || snapshot.inFlight != 4 || snapshot.queued != 3 {
t.Fatalf("snapshot=%+v", snapshot)
}
t.Run("fetches selected aggregate", func(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "application/json")
_, _ = io.WriteString(w, validStatus)
}))
defer server.Close()
got, fetchErr := fetchDevRepeatGuardCapacityStatus(
server.Client(), server.URL, providerIDs,
)
if fetchErr != nil {
t.Fatal(fetchErr)
}
if got != snapshot {
t.Fatalf("fetched snapshot=%+v, want %+v", got, snapshot)
}
})
t.Run("observes peak queue and final recovery", func(t *testing.T) {
var requestCount atomic.Int32
snapshots := []struct {
inFlight int
queued int
}{
{inFlight: 2},
{inFlight: 4, queued: 1},
{inFlight: 1},
{},
}
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
index := int(requestCount.Add(1)) - 1
if index >= len(snapshots) {
index = len(snapshots) - 1
}
current := snapshots[index]
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(map[string]any{
"nodes": []any{map[string]any{
"provider_snapshots": []any{map[string]any{
"id": "provider-a", "capacity": 4,
"in_flight": current.inFlight, "queued": current.queued,
}},
}},
})
}))
defer server.Close()
doneC := make(chan struct{})
close(doneC)
evidence, observeErr := observeDevRepeatGuardCapacityRun(
server.Client(), server.URL, []string{"provider-a"}, 1,
doneC, make(chan struct{}), 2*time.Second,
)
if observeErr != nil {
t.Fatal(observeErr)
}
if evidence.ConfiguredCapacity != 4 ||
evidence.PeakInFlight != 4 ||
evidence.PeakQueued != 1 ||
evidence.FinalInFlight != 0 ||
evidence.FinalQueued != 0 {
t.Fatalf("capacity evidence=%+v", evidence)
}
})
for _, tc := range []struct {
name string
raw string
}{
{name: "malformed JSON", raw: `{"nodes":[`},
{
name: "duplicate selected provider",
raw: `{"nodes":[{"provider_snapshots":[
{"id":"provider-a","capacity":2,"in_flight":0,"queued":0},
{"id":"provider-a","capacity":2,"in_flight":0,"queued":0},
{"id":"provider-b","capacity":2,"in_flight":0,"queued":0}
]}]}`,
},
{
name: "missing selected provider",
raw: `{"nodes":[{"provider_snapshots":[
{"id":"provider-a","capacity":2,"in_flight":0,"queued":0}
]}]}`,
},
{
name: "negative selected counter",
raw: `{"nodes":[{"provider_snapshots":[
{"id":"provider-a","capacity":2,"in_flight":-1,"queued":0},
{"id":"provider-b","capacity":2,"in_flight":0,"queued":0}
]}]}`,
},
} {
t.Run(tc.name, func(t *testing.T) {
if _, parseErr := parseDevRepeatGuardCapacityStatus(
[]byte(tc.raw), providerIDs,
); parseErr == nil {
t.Fatal("expected capacity status validation failure")
}
})
}
for _, rawIDs := range []string{"", "provider-a,", "provider-a,provider-a"} {
t.Run("rejects provider IDs "+strconv.Quote(rawIDs), func(t *testing.T) {
if _, parseErr := parseDevRepeatGuardProviderIDs(rawIDs); parseErr == nil {
t.Fatal("expected provider ID validation failure")
}
})
}
validEvidence := devRepeatGuardCapacityEvidence{
Run: 1, ConfiguredCapacity: 4, PeakInFlight: 4, PeakQueued: 1,
FinalInFlight: 0, FinalQueued: 0,
}
if err := validateDevRepeatGuardCapacityEvidence(validEvidence); err != nil {
t.Fatal(err)
}
for _, tc := range []struct {
name string
mutate func(*devRepeatGuardCapacityEvidence)
}{
{name: "wrong capacity", mutate: func(value *devRepeatGuardCapacityEvidence) {
value.ConfiguredCapacity = 3
}},
{name: "missing in flight peak", mutate: func(value *devRepeatGuardCapacityEvidence) {
value.PeakInFlight = 3
}},
{name: "missing queue pressure", mutate: func(value *devRepeatGuardCapacityEvidence) {
value.PeakQueued = 0
}},
{name: "not recovered", mutate: func(value *devRepeatGuardCapacityEvidence) {
value.FinalQueued = 1
}},
} {
t.Run(tc.name, func(t *testing.T) {
value := validEvidence
tc.mutate(&value)
if validateErr := validateDevRepeatGuardCapacityEvidence(value); validateErr == nil {
t.Fatal("expected capacity evidence validation failure")
}
})
}
}
// --- S16: host recovery preparer lifecycle ----------------------------------
// recordingPreparer is a test-only host RecoveryPlanPreparer. Production OpenAI
// surfaces register no preparer; this fixture exercises the host wiring of the
// optional one-shot preparation seam (S23): it must run at most once per plan,
// after attempt ownership is closed and before the rebuild/dispatch.
type recordingPreparer struct {
mu sync.Mutex
calls int
failWith error
directive streamgate.RecoveryDirective
}
func (p *recordingPreparer) PrepareRecoveryPlan(ctx context.Context, plan streamgate.RecoveryPlan, snapshot streamgate.RecoveryPreparationSnapshot) (streamgate.RecoveryDirective, error) {
p.mu.Lock()
p.calls++
failWith := p.failWith
directive := p.directive
p.mu.Unlock()
if failWith != nil {
return streamgate.RecoveryDirective{}, failWith
}
return directive, nil
}
func (p *recordingPreparer) callCount() int {
p.mu.Lock()
defer p.mu.Unlock()
return p.calls
}
// recordingPreparationSnapshot is the minimal bounded snapshot seam; it only
// records that Core released it exactly once.
type recordingPreparationSnapshot struct {
ref string
releases *int32
}
func (s *recordingPreparationSnapshot) SnapshotRef() string { return s.ref }
func (s *recordingPreparationSnapshot) Release() error {
atomic.AddInt32(s.releases, 1)
return nil
}
type recordingPreparationSnapshotFactory struct {
ref string
creates *int32
releases *int32
}
func (f *recordingPreparationSnapshotFactory) CreatePreparationSnapshot(ctx context.Context, snapshot streamgate.RequestRuntimeSnapshot, binding streamgate.AttemptBinding) (streamgate.RecoveryPreparationSnapshot, error) {
atomic.AddInt32(f.creates, 1)
return &recordingPreparationSnapshot{ref: f.ref, releases: f.releases}, nil
}
// TestStreamGateChatRecoveryPreparerLifecycle pins the host preparer contract
// for one recovery cycle: on success the prepared directive is what gets
// rebuilt and dispatched exactly once, and on failure the cycle stops before
// any dispatch and converges on a single error terminal with no leaked content.
func TestStreamGateChatRecoveryPreparerLifecycle(t *testing.T) {
for _, tc := range []struct {
name string
preparerFail bool
wantDispatch int
}{
{"preparer success dispatches once", false, 1},
{"preparer failure dispatches nothing", true, 0},
} {
tc := tc
t.Run(tc.name, func(t *testing.T) {
srv, dc, initial, fake := buildStreamGateChatFixture(t, true, 3)
fake.eventRuns = []chan *iop.RunEvent{bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "final answer"},
&iop.RunEvent{Type: "complete"},
)}
fake.runIDs = []string{"run-attempt-2"}
snapRef, err := dc.ingress.recoveryRef()
if err != nil {
t.Fatalf("recoveryRef: %v", err)
}
directive, err := streamgate.NewRecoveryDirectiveExact(snapRef.SnapshotRef())
if err != nil {
t.Fatalf("NewRecoveryDirectiveExact: %v", err)
}
preparer := &recordingPreparer{directive: directive}
if tc.preparerFail {
preparer.failWith = fmt.Errorf("injected preparer failure")
}
var creates, releases int32
factory := &recordingPreparationSnapshotFactory{ref: "openai.prep.1", creates: &creates, releases: &releases}
registry := streamGateTestRegistry(t, newInjectedViolationRegistration(t, dc, "test.injected_violation", 1))
w := newRecordingResponseWriter()
sink := newOpenAIChatSSEReleaseSink(w, w, "chatcmpl-run-attempt-1", time.Now().Unix(), "client-model")
rt, _, err := srv.buildOpenAIChatStreamGateRuntimeFor(dc, openAIChatStreamGateConfig{
mode: openAIChatGateModeLive,
initial: openAIAttemptTransport{path: openAIAdmissionRun, run: initial},
dispatch: initial.Dispatch(),
closeAll: initial.Close,
sink: sink,
selector: newOpenAIStreamGateCodecSelector(openAIStreamGateCodecNormalized),
registry: registry,
preparer: preparer,
prepFactory: factory,
})
if err != nil {
t.Fatalf("buildOpenAIChatStreamGateRuntimeFor: %v", err)
}
if err := rt.Run(context.Background()); err != nil {
t.Fatalf("rt.Run: %v", err)
}
if got := preparer.callCount(); got != 1 {
t.Fatalf("preparer calls: got %d, want exactly 1 per recovery plan", got)
}
if got := atomic.LoadInt32(&creates); got != 1 {
t.Fatalf("preparation snapshot creations: got %d, want 1", got)
}
if got := atomic.LoadInt32(&releases); got != 1 {
t.Fatalf("preparation snapshot releases: got %d, want exactly 1", got)
}
if got := len(fake.reqsSnapshot()); got != tc.wantDispatch {
t.Fatalf("recovery dispatches: got %d, want %d", got, tc.wantDispatch)
}
body := w.body.String()
if strings.Contains(body, "should-be-discarded") {
t.Fatalf("the aborted attempt leaked into the response: %s", body)
}
if tc.preparerFail {
if !strings.Contains(body, `"error"`) {
t.Fatalf("a failed preparation must terminate with an error, body=%s", body)
}
return
}
if got := joinedContent(parseSSEChatChunks(t, body)); got != "final answer" {
t.Fatalf("joined content: got %q, want the prepared recovery attempt's content", got)
}
})
}
}
// --- S16: producer backpressure while an evaluator holds --------------------
// blockingEvaluationFilter blocks inside Evaluate until it is released, so a
// fixture can observe that the unbuffered provider producer cannot advance
// while the Core is evaluating an epoch.
type blockingEvaluationFilter struct {
streamgate.FilterBase
entered chan struct{}
release chan struct{}
once sync.Once
}
func newBlockingEvaluationFilter(t *testing.T, id string) *blockingEvaluationFilter {
t.Helper()
base, err := streamgate.NewFilterBase(id)
if err != nil {
t.Fatalf("NewFilterBase: %v", err)
}
return &blockingEvaluationFilter{
FilterBase: base,
entered: make(chan struct{}, 1),
release: make(chan struct{}),
}
}
func (f *blockingEvaluationFilter) Applies(streamgate.FilterContext) bool { return true }
func (f *blockingEvaluationFilter) HoldRequirement(streamgate.FilterContext) streamgate.FilterHoldRequirement {
req, err := streamgate.NewFilterHoldRequirementRolling(streamGateChannelDefault, []streamgate.EventKind{streamgate.EventKindTextDelta}, 1)
if err != nil {
panic(err)
}
return req
}
func (f *blockingEvaluationFilter) Evaluate(ctx context.Context, fctx streamgate.FilterContext, batch streamgate.EvidenceBatch) (streamgate.FilterDecision, error) {
f.once.Do(func() {
f.entered <- struct{}{}
select {
case <-f.release:
case <-ctx.Done():
}
})
kind := streamgate.EventKindTextDelta
if events := batch.Events(); len(events) > 0 {
kind = events[len(events)-1].Kind()
}
evidence, err := streamgate.NewSanitizedEvidence(kind, streamGateChannelDefault, "test.rule", "test_blocking_eval", streamgate.FixedFingerprint{0x03}, 1, 0, streamgate.FilterOutcomeKindEvaluated, batch.CapturedAt())
if err != nil {
return streamgate.FilterDecision{}, err
}
return streamgate.NewFilterDecision(streamgate.FilterDecisionKindPass, "test.consumer", f.ID(), "test.rule", evidence, nil)
}
var _ streamgate.Filter = (*blockingEvaluationFilter)(nil)
// TestStreamGateChatProducerBackpressureDuringEvaluation verifies the Core
// applies real backpressure: while a blocking filter is still evaluating the
// first epoch, the unbuffered provider producer cannot complete its next send,
// and it resumes only once the evaluation is released.
func TestStreamGateChatProducerBackpressureDuringEvaluation(t *testing.T) {
srv, dc, initial, _ := buildStreamGateChatFixture(t, true, 3)
events := make(chan *iop.RunEvent) // unbuffered: the producer blocks unless the source is reading
initial.events = events
sent := make(chan int, 3)
go func() {
for i, ev := range []*iop.RunEvent{
{Type: "delta", Delta: "one "},
{Type: "delta", Delta: "two"},
{Type: "complete"},
} {
events <- ev
sent <- i + 1
}
}()
blocking := newBlockingEvaluationFilter(t, "test.blocking_eval")
reg, err := streamgate.NewFilterRegistration(blocking, streamGateNoopCapability, true, streamgate.FilterEnforcementBlocking, streamGateFilterTimeout, 5)
if err != nil {
t.Fatalf("NewFilterRegistration: %v", err)
}
registry := streamGateTestRegistry(t, reg)
w := newRecordingResponseWriter()
rt, _, err := srv.buildOpenAIChatStreamGateRuntime(dc, initial, newOpenAIChatSSEReleaseSink(w, w, "chatcmpl-run-attempt-1", time.Now().Unix(), "client-model"), registry)
if err != nil {
t.Fatalf("buildOpenAIChatStreamGateRuntime: %v", err)
}
done := make(chan error, 1)
go func() { done <- rt.Run(context.Background()) }()
select {
case <-blocking.entered:
case <-time.After(5 * time.Second):
t.Fatal("the blocking evaluator was never entered")
}
// The first event was consumed by the event source; the second send must be
// parked because the runtime is inside the evaluator and not reading.
select {
case got := <-sent:
if got != 1 {
t.Fatalf("first completed send: got %d, want 1", got)
}
case <-time.After(5 * time.Second):
t.Fatal("the first send never completed")
}
select {
case got := <-sent:
t.Fatalf("send %d completed while the evaluator still held the epoch: producer backpressure is missing", got)
case <-time.After(150 * time.Millisecond):
}
close(blocking.release)
select {
case runErr := <-done:
if runErr != nil {
t.Fatalf("rt.Run: %v", runErr)
}
case <-time.After(10 * time.Second):
t.Fatal("rt.Run did not finish after the evaluation was released")
}
for want := 2; want <= 3; want++ {
select {
case got := <-sent:
if got != want {
t.Fatalf("resumed send: got %d, want %d", got, want)
}
case <-time.After(5 * time.Second):
t.Fatalf("send %d never completed after the evaluation was released", want)
}
}
if got := joinedContent(parseSSEChatChunks(t, w.body.String())); got != "one two" {
t.Fatalf("joined content: got %q, want %q", got, "one two")
}
}
// TestOpenAIStreamGateRecoveryCandidateRejectionIsBadRequest proves that a
// provider-pool capability rejection during Core recovery keeps the same public
// admission contract as initial and queued rejection. The rejected recovery
// performs no provider transport dispatch and commits a single pre-stream 400.
func TestOpenAIStreamGateRecoveryCandidateRejectionIsBadRequest(t *testing.T) {
service := newScriptedPoolRunService(
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathNormalized), runID: "recovery-reject-1",
provider: "prov-a", target: "served-a",
runEvents: bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "discarded"},
&iop.RunEvent{Type: "complete"},
),
},
scriptedPoolAttempt{err: edgeservice.ErrProviderPoolCandidateRejected},
)
rawBody := []byte(`{"model":"client-model","messages":[{"role":"user","content":"hi"}],"stream":true}`)
srv, dc := buildStreamGatePoolChatFixture(t, service, rawBody, 3)
w, sink, runErr := runPoolChatStreamGate(t, srv, dc, newInjectedViolationRegistration(t, dc, "test.injected_violation", 1))
if runErr != nil {
t.Fatalf("runtime.Run: %v", runErr)
}
if w.code != http.StatusBadRequest {
t.Fatalf("status: got %d, want %d; body=%q", w.code, http.StatusBadRequest, w.body.String())
}
if !strings.Contains(w.body.String(), openAIStreamGateCandidateRejectedMessage) {
t.Fatalf("candidate rejection message missing: %q", w.body.String())
}
if strings.Contains(w.body.String(), "discarded") {
t.Fatalf("rejected attempt content leaked: %q", w.body.String())
}
committed, success := sink.terminalStatus()
if !committed || success {
t.Fatalf("terminal status: got (%v,%v), want (true,false)", committed, success)
}
pools, _, _, _, runReqs, tunnelReqs := service.snapshot()
if pools != 2 {
t.Fatalf("pool admissions: got %d, want initial + one rejected recovery", pools)
}
if len(runReqs) != 1 || len(tunnelReqs) != 0 {
t.Fatalf("provider transports: normalized=%d tunnel=%d, want only the initial normalized dispatch", len(runReqs), len(tunnelReqs))
}
}
// --- S16: Responses provider-pool tunnel recovery ---------------------------
// buildStreamGateResponsesPoolFixture builds a streaming /v1/responses
// provider-pool passthrough request bound to a scripted pool double, mirroring
// what handleResponsesProviderPool freezes before its first SubmitProviderPool.
func buildStreamGateResponsesPoolFixture(t *testing.T, service runService, maxRequestFaultRecovery int) (*Server, *responsesRequestContext, edgeservice.ProviderPoolDispatchRequest) {
t.Helper()
rawBody := []byte(`{"model":"client-model","input":"CALLER_RESPONSES_INPUT_MUST_NOT_BE_COPIED","stream":true}`)
fault := maxRequestFaultRecovery
srv := NewServer(config.EdgeOpenAIConf{
Adapter: "openai-compat",
Target: "served-model",
TimeoutSec: 15,
StreamEvidenceGate: config.StreamEvidenceGateConf{
Enabled: true,
MaxRequestFaultRecovery: &fault,
},
}, service, nil)
srv.SetModelCatalog([]config.ModelCatalogEntry{{ID: "client-model", ContextWindowTokens: 4096}})
route := routeDispatch{ProviderPool: true, TimeoutSec: 15}
base := newTestRequestContext(t, route, rawBody)
base.endpoint = usageEndpointResponses
base.callerMetadata = map[string]string{"principal_ref": "principal-test"}
base.estimate = 73
base.contextClass = "long"
requestCtx := &responsesRequestContext{
openAIRequestContext: base,
envelope: responsesEnvelope{Model: "client-model", Stream: true},
}
metadata := map[string]string{
"principal_ref": "principal-test",
"openai_model": "client-model",
"openai_stream": "true",
"estimated_input_tokens": "73",
"context_class": "long",
}
poolReq := edgeservice.ProviderPoolDispatchRequest{
Run: edgeservice.SubmitRunRequest{
ModelGroupKey: "client-model",
Metadata: metadata,
ProviderPool: true,
},
Tunnel: edgeservice.SubmitProviderTunnelRequest{
ModelGroupKey: "client-model",
Method: http.MethodPost,
Path: "/v1/responses",
Stream: true,
TimeoutSec: 15,
Metadata: metadata,
ProviderPool: true,
},
PrepareTunnel: func(tunnelReq edgeservice.SubmitProviderTunnelRequest) (edgeservice.SubmitProviderTunnelRequest, error) {
return tunnelReq, nil
},
PrepareRun: func(runReq edgeservice.SubmitRunRequest) (edgeservice.SubmitRunRequest, error) {
return runReq, nil
},
}
poolReq.Tunnel.BuildBody = func(target string) ([]byte, error) {
body, err := requestCtx.ingress.canonicalBody()
if err != nil {
return nil, err
}
return rewriteResponsesModel(body, target)
}
return srv, requestCtx, poolReq
}
func responsesTunnelFrames(payload string) chan *iop.ProviderTunnelFrame {
return bufferedTunnelFrames(
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_RESPONSE_START, StatusCode: 200, Headers: map[string]string{"Content-Type": "text/event-stream"}},
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: []byte(payload)},
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_END, End: true},
)
}
func responsesStreamPrefix(responseID, itemID, text string) string {
return fmt.Sprintf("data: {\"type\":\"response.created\",\"response\":{\"id\":\"%s\",\"object\":\"response\",\"status\":\"in_progress\"},\"sequence_number\":1}\n\n"+
"data: {\"type\":\"response.output_item.added\",\"output_index\":0,\"item\":{\"id\":\"%s\",\"type\":\"message\",\"status\":\"in_progress\",\"role\":\"assistant\",\"content\":[]},\"sequence_number\":2}\n\n"+
"data: {\"type\":\"response.content_part.added\",\"item_id\":\"%s\",\"output_index\":0,\"content_index\":0,\"part\":{\"type\":\"output_text\",\"text\":\"\",\"annotations\":[],\"logprobs\":[]},\"sequence_number\":3}\n\n"+
"data: {\"type\":\"response.output_text.delta\",\"item_id\":\"%s\",\"output_index\":0,\"content_index\":0,\"delta\":\"%s\",\"logprobs\":[],\"sequence_number\":4}\n\n", responseID, itemID, itemID, itemID, text)
}
func responsesStreamTerminal(responseID, itemID, text string, sequence int) string {
return fmt.Sprintf("data: {\"type\":\"response.output_text.done\",\"item_id\":\"%s\",\"output_index\":0,\"content_index\":0,\"text\":\"%s\",\"logprobs\":[],\"sequence_number\":%d}\n\n"+
"data: {\"type\":\"response.content_part.done\",\"item_id\":\"%s\",\"output_index\":0,\"content_index\":0,\"part\":{\"type\":\"output_text\",\"text\":\"%s\",\"annotations\":[],\"logprobs\":[]},\"sequence_number\":%d}\n\n"+
"data: {\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{\"id\":\"%s\",\"type\":\"message\",\"status\":\"completed\",\"role\":\"assistant\",\"content\":[{\"type\":\"output_text\",\"text\":\"%s\",\"annotations\":[],\"logprobs\":[]}]},\"sequence_number\":%d}\n\n"+
"data: {\"type\":\"response.completed\",\"response\":{\"id\":\"%s\",\"object\":\"response\",\"status\":\"completed\",\"output\":[{\"id\":\"%s\",\"type\":\"message\",\"status\":\"completed\",\"role\":\"assistant\",\"content\":[{\"type\":\"output_text\",\"text\":\"%s\",\"annotations\":[],\"logprobs\":[]}]}],\"usage\":{\"input_tokens\":1,\"output_tokens\":1,\"total_tokens\":2}},\"sequence_number\":%d}\n\n"+
"data: [DONE]\n\n", itemID, text, sequence, itemID, text, sequence+1, itemID, text, sequence+2, responseID, itemID, text, sequence+3)
}
func responsesTextDelta(itemID, text string, sequence int) string {
return fmt.Sprintf("data: {\"type\":\"response.output_text.delta\",\"item_id\":\"%s\",\"output_index\":0,\"content_index\":0,\"delta\":\"%s\",\"sequence_number\":%d}\n\n", itemID, text, sequence)
}
func responsesRecoveryTunnelFrames(text string, inputTokens, outputTokens int) chan *iop.ProviderTunnelFrame {
return bufferedTunnelFrames(
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_RESPONSE_START, StatusCode: 200, Headers: map[string]string{"Content-Type": "application/json"}},
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: []byte(fmt.Sprintf("{\"type\":\"response.output_text.delta\",\"delta\":\"%s\"}", text))},
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_USAGE, Usage: &iop.Usage{InputTokens: int32(inputTokens), OutputTokens: int32(outputTokens)}},
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_END, End: true},
)
}
// TestStreamGateResponsesPoolRecoveryTerminal proves the real Responses Core
// transaction from an initial provider tunnel through one continuation
// readmission. The replacement may stay on the tunnel path or switch to a
// normalized RunEvent path; both must commit only the replacement attempt.
func TestStreamGateResponsesPoolRecoveryTerminal(t *testing.T) {
for _, tc := range []struct {
name string
path string
wantCodec openAIStreamGateCodec
wantOutput string
wantUsage usageObservation
}{
{
name: "tunnel replacement",
path: string(edgeservice.ProviderPoolPathTunnel),
wantCodec: openAIStreamGateCodecTunnel,
wantOutput: "recovered tunnel",
wantUsage: usageObservation{inputTokens: 31, outputTokens: 17},
},
{
name: "normalized replacement",
path: string(edgeservice.ProviderPoolPathNormalized),
wantCodec: openAIStreamGateCodecNormalized,
wantOutput: "recovered normalized",
wantUsage: usageObservation{inputTokens: 29, outputTokens: 13},
},
} {
tc := tc
t.Run(tc.name, func(t *testing.T) {
initialPayload := responsesStreamPrefix("resp-initial", "msg-initial", "safe prefix ") +
"data: {\"type\":\"response.output_text.delta\",\"item_id\":\"msg-initial\",\"output_index\":0,\"content_index\":0,\"delta\":\"discarded repeat tail\",\"sequence_number\":5}\n\n"
replacement := scriptedPoolAttempt{
path: tc.path, runID: "resp-attempt-2", provider: "prov-b", target: "served-b",
}
if tc.path == string(edgeservice.ProviderPoolPathTunnel) {
replacement.frames = responsesRecoveryTunnelFrames(tc.wantOutput, tc.wantUsage.inputTokens, tc.wantUsage.outputTokens)
} else {
replacement.runEvents = bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: tc.wantOutput},
&iop.RunEvent{Type: "complete", Usage: &iop.Usage{InputTokens: int32(tc.wantUsage.inputTokens), OutputTokens: int32(tc.wantUsage.outputTokens)}},
)
}
service := newScriptedPoolRunService(
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathTunnel), runID: "resp-attempt-1",
provider: "prov-a", target: "served-a",
frames: responsesTunnelFrames(initialPayload),
},
replacement,
)
srv, requestCtx, poolReq := buildStreamGateResponsesPoolFixture(t, service, 3)
result, err := service.SubmitProviderPool(context.Background(), poolReq)
if err != nil {
t.Fatalf("initial SubmitProviderPool: %v", err)
}
dc := newOpenAIResponsesPoolTunnelDispatchContext(requestCtx, poolReq)
snapRef, err := requestCtx.ingress.recoveryRef()
if err != nil {
t.Fatalf("recoveryRef: %v", err)
}
violation := newInjectedContinuationViolationFilter(t, "test.responses.resume", 1, snapRef.SnapshotRef(), len("safe prefix "))
reg, err := streamgate.NewFilterRegistration(violation, streamGateNoopCapability, true, streamgate.FilterEnforcementBlocking, streamGateFilterTimeout, injectedViolationPriority)
if err != nil {
t.Fatalf("NewFilterRegistration: %v", err)
}
w := newRecordingResponseWriter()
holder := &openAIResponsesResultHolder{}
selector := newOpenAIStreamGateCodecSelector(openAIStreamGateCodecTunnel)
sink := newOpenAIResponsesPoolReleaseSink(w, holder, selector)
runtime, usage, err := srv.buildOpenAIResponsesStreamGateRuntimeFromAttempt(
dc,
openAIAttemptTransport{path: openAIAdmissionTunnel, tunnel: result.Tunnel},
result.Tunnel.Dispatch(),
result.Tunnel.Close,
sink,
streamGateTestRegistry(t, reg),
)
if err != nil {
t.Fatalf("buildOpenAIResponsesStreamGateRuntimeFromAttempt: %v", err)
}
if err := runtime.Run(context.Background()); err != nil {
t.Fatalf("runtime.Run: %v", err)
}
if err := runtime.CloseRequestResources(context.Background(), true); err != nil {
t.Fatalf("CloseRequestResources: %v", err)
}
if got := service.poolSubmits(); got != 2 {
t.Fatalf("SubmitProviderPool calls: got %d, want initial + one recovery", got)
}
if got := selector.get(); got != tc.wantCodec {
t.Fatalf("resolved codec = %q, want %q", got, tc.wantCodec)
}
committed, success := sink.terminalStatus()
if !committed || !success {
t.Fatalf("terminal status = (%v,%v), want (true,true)", committed, success)
}
body := w.body.String()
payloads := assertResponsesSSELifecycle(t, body, "response.completed")
var terminal map[string]any
for _, payload := range payloads {
if payload["type"] == "response.completed" {
terminal, _ = payload["response"].(map[string]any)
break
}
}
if got, _ := terminal["model"].(string); got != "served-b" {
t.Fatalf("response.completed model = %q, want selected replacement model served-b: %#v", got, terminal)
}
terminalUsage, _ := terminal["usage"].(map[string]any)
if got := terminalUsage["input_tokens"]; got != float64(tc.wantUsage.inputTokens) {
t.Fatalf("response.completed input_tokens = %#v, want %d: %#v", got, tc.wantUsage.inputTokens, terminalUsage)
}
if got := terminalUsage["output_tokens"]; got != float64(tc.wantUsage.outputTokens) {
t.Fatalf("response.completed output_tokens = %#v, want %d: %#v", got, tc.wantUsage.outputTokens, terminalUsage)
}
if got := terminalUsage["total_tokens"]; got != float64(tc.wantUsage.inputTokens+tc.wantUsage.outputTokens) {
t.Fatalf("response.completed total_tokens = %#v, want %d: %#v", got, tc.wantUsage.inputTokens+tc.wantUsage.outputTokens, terminalUsage)
}
if got := responsesSSEOutputText(payloads); got != "safe prefix "+tc.wantOutput {
t.Fatalf("Responses SSE output = %q, want safe prefix plus replacement; body=%q", got, body)
}
if strings.Contains(body, "discarded repeat tail") || strings.Contains(body, "CALLER_RESPONSES_INPUT_MUST_NOT_BE_COPIED") {
t.Fatalf("response body did not preserve only the safe prefix and replacement: %q", body)
}
if got := w.Header().Get("Content-Type"); got != "text/event-stream" {
t.Fatalf("Content-Type = %q, want text/event-stream", got)
}
if got := w.headerCallCount(); got != 1 {
t.Fatalf("WriteHeader call count: got %d, want one terminal commit", got)
}
if got := usage.get(); got.inputTokens != tc.wantUsage.inputTokens || got.outputTokens != tc.wantUsage.outputTokens {
t.Fatalf("terminal usage = %#v, want %#v", got, tc.wantUsage)
}
pools, cancels, _, _, runRequests, tunnelRequests := service.snapshot()
if pools != 2 || countString(cancels, "resp-attempt-1") != 1 {
t.Fatalf("recovery lifecycle pools=%d cancels=%v, want two admissions and one initial abort", pools, cancels)
}
if tc.path == string(edgeservice.ProviderPoolPathNormalized) {
if len(runRequests) != 1 || runRequests[0].Metadata["openai_stream"] != "false" || runRequests[0].EstimatedInputTokens != 68 || runRequests[0].ContextClass != "normal" {
t.Fatalf("normalized replacement request = %#v", runRequests)
}
if strings.Contains(runRequests[0].Prompt, "CALLER_RESPONSES_INPUT_MUST_NOT_BE_COPIED") || runRequests[0].Input["responses_resume_content"] != "safe prefix " {
t.Fatalf("normalized replacement did not preserve private resume provenance: %#v", runRequests[0])
}
return
}
if len(tunnelRequests) != 2 {
t.Fatalf("tunnel requests = %d, want initial plus replacement", len(tunnelRequests))
}
resumeTunnel := tunnelRequests[1]
if resumeTunnel.Stream || resumeTunnel.Metadata["openai_stream"] != "false" || resumeTunnel.EstimatedInputTokens != 68 || resumeTunnel.ContextClass != "normal" {
t.Fatalf("tunnel replacement admission did not use the rebuilt attempt: %#v", resumeTunnel)
}
resumeBody, err := resumeTunnel.BuildBody("served-b")
if err != nil {
t.Fatalf("replacement BuildBody: %v", err)
}
if strings.Contains(string(resumeBody), "CALLER_RESPONSES_INPUT_MUST_NOT_BE_COPIED") || !strings.Contains(string(resumeBody), "safe prefix ") || !strings.Contains(string(resumeBody), "Continue from the provided content and reasoning output") {
t.Fatalf("tunnel replacement body does not contain only the private resume request: %s", resumeBody)
}
})
}
}
func TestOpenAIResponsesPoolSyntheticLifecycle(t *testing.T) {
w := newRecordingResponseWriter()
sink := newOpenAIResponsesPoolReleaseSink(w, &openAIResponsesResultHolder{}, newOpenAIStreamGateCodecSelector(openAIStreamGateCodecNormalized))
sink.bindAttempt(false, edgeservice.RunDispatch{Target: "served-synthetic"})
usage := &openAIStreamGateUsageHolder{}
usage.set(usageObservation{inputTokens: 7, outputTokens: 11})
sink.setUsageHolder(usage)
start, err := streamgate.NewResponseStart(streamGateChannelDefault, http.StatusOK, map[string]string{"Content-Type": "application/json"}, time.Now())
if err != nil {
t.Fatalf("NewResponseStart: %v", err)
}
if _, err := sink.CommitResponseStart(context.Background(), start); err != nil {
t.Fatalf("CommitResponseStart: %v", err)
}
for _, event := range []streamgate.ReleaseEvent{
mustReleaseText(t, "answer"),
mustReleaseReasoning(t, "because"),
mustReleaseToolCall(t, "call-1", "lookup", `{"q":"IOP"}`),
} {
if _, err := sink.Release(context.Background(), event); err != nil {
t.Fatalf("Release(%q): %v", event.Kind(), err)
}
}
terminal, err := streamgate.NewSuccessTerminalResult(streamGateChannelDefault, time.Now())
if err != nil {
t.Fatalf("NewSuccessTerminalResult: %v", err)
}
if _, err := sink.CommitTerminal(context.Background(), terminal); err != nil {
t.Fatalf("CommitTerminal: %v", err)
}
payloads := assertResponsesSSELifecycle(t, w.body.String(), "response.completed")
var kinds []string
for _, payload := range payloads {
kind, _ := payload["type"].(string)
kinds = append(kinds, kind)
}
wantKinds := []string{
"response.created", "response.output_item.added", "response.content_part.added", "response.output_text.delta",
"response.output_item.added", "response.content_part.added", "response.reasoning_text.delta",
"response.output_item.added", "response.function_call_arguments.delta",
"response.output_text.done", "response.content_part.done", "response.output_item.done",
"response.reasoning_text.done", "response.content_part.done", "response.output_item.done",
"response.function_call_arguments.done", "response.output_item.done", "response.completed",
}
if !reflect.DeepEqual(kinds, wantKinds) {
t.Fatalf("synthetic lifecycle kinds = %#v, want %#v", kinds, wantKinds)
}
completed := payloads[len(payloads)-1]["response"].(map[string]any)
if got := completed["model"]; got != "served-synthetic" {
t.Fatalf("synthetic terminal model = %#v", got)
}
if got := completed["output_text"]; got != "answer" {
t.Fatalf("synthetic terminal output_text = %#v", got)
}
terminalUsage := completed["usage"].(map[string]any)
if terminalUsage["input_tokens"] != float64(7) || terminalUsage["output_tokens"] != float64(11) || terminalUsage["total_tokens"] != float64(18) {
t.Fatalf("synthetic terminal usage = %#v", terminalUsage)
}
}
func mustReleaseText(t *testing.T, value string) streamgate.ReleaseEvent {
t.Helper()
event, err := streamgate.NewReleaseTextDeltaEvent(streamGateChannelDefault, value, time.Now())
if err != nil {
t.Fatalf("NewReleaseTextDeltaEvent: %v", err)
}
return event
}
func mustReleaseReasoning(t *testing.T, value string) streamgate.ReleaseEvent {
t.Helper()
event, err := streamgate.NewReleaseReasoningDeltaEvent(streamGateChannelDefault, value, time.Now())
if err != nil {
t.Fatalf("NewReleaseReasoningDeltaEvent: %v", err)
}
return event
}
func mustReleaseToolCall(t *testing.T, id, name, arguments string) streamgate.ReleaseEvent {
t.Helper()
event, err := streamgate.NewReleaseToolCallFragmentEvent(streamGateChannelDefault, id, name, arguments, time.Now())
if err != nil {
t.Fatalf("NewReleaseToolCallFragmentEvent: %v", err)
}
return event
}
// rawResponsesItemSpec describes a raw provider item the prefix builder opens.
type rawResponsesItemSpec struct {
kind string // "message", "reasoning", or "function"
responseID string
itemID string
outputIndex int
callID string
name string
value string // content text for message/reasoning, arguments for function
}
func responsesRawFrame(seq *int, payload map[string]any) string {
*seq++
payload["sequence_number"] = *seq
encoded, err := json.Marshal(payload)
if err != nil {
panic(err)
}
return fmt.Sprintf("data: %s\n\n", encoded)
}
// buildRawResponsesItemPrefix emits a schema-valid raw provider SSE prefix that
// opens one item and streams it up to boundary ("delta", "value_done",
// "content_done", or "item_done"). "content_done" is only valid for content
// items; a function item has no content part.
func buildRawResponsesItemPrefix(spec rawResponsesItemSpec, boundary string) string {
seq := 0
var b strings.Builder
b.WriteString(responsesRawFrame(&seq, map[string]any{"type": "response.created", "response": map[string]any{"id": spec.responseID, "object": "response", "status": "in_progress"}}))
if spec.kind == "function" {
b.WriteString(responsesRawFrame(&seq, map[string]any{"type": "response.output_item.added", "output_index": spec.outputIndex, "item": map[string]any{"id": spec.itemID, "type": "function_call", "status": "in_progress", "call_id": spec.callID, "name": spec.name, "arguments": ""}}))
b.WriteString(responsesRawFrame(&seq, map[string]any{"type": "response.function_call_arguments.delta", "item_id": spec.itemID, "output_index": spec.outputIndex, "delta": spec.value}))
if boundary == "delta" {
return b.String()
}
b.WriteString(responsesRawFrame(&seq, map[string]any{"type": "response.function_call_arguments.done", "item_id": spec.itemID, "output_index": spec.outputIndex, "name": spec.name, "arguments": spec.value}))
if boundary == "value_done" {
return b.String()
}
b.WriteString(responsesRawFrame(&seq, map[string]any{"type": "response.output_item.done", "output_index": spec.outputIndex, "item": map[string]any{"id": spec.itemID, "type": "function_call", "status": "completed", "call_id": spec.callID, "name": spec.name, "arguments": spec.value}}))
return b.String()
}
partType := "output_text"
deltaType := "response.output_text.delta"
doneType := "response.output_text.done"
if spec.kind == "reasoning" {
partType, deltaType, doneType = "reasoning_text", "response.reasoning_text.delta", "response.reasoning_text.done"
}
openingItem := map[string]any{"id": spec.itemID, "type": spec.kind, "status": "in_progress", "content": []any{}}
openPart := map[string]any{"type": partType, "text": ""}
donePart := map[string]any{"type": partType, "text": spec.value}
if spec.kind == "message" {
openingItem["role"] = "assistant"
openPart["annotations"], openPart["logprobs"] = []any{}, []any{}
donePart["annotations"], donePart["logprobs"] = []any{}, []any{}
}
b.WriteString(responsesRawFrame(&seq, map[string]any{"type": "response.output_item.added", "output_index": spec.outputIndex, "item": openingItem}))
b.WriteString(responsesRawFrame(&seq, map[string]any{"type": "response.content_part.added", "item_id": spec.itemID, "output_index": spec.outputIndex, "content_index": 0, "part": openPart}))
deltaPayload := map[string]any{"type": deltaType, "item_id": spec.itemID, "output_index": spec.outputIndex, "content_index": 0, "delta": spec.value}
if spec.kind == "message" {
deltaPayload["logprobs"] = []any{}
}
b.WriteString(responsesRawFrame(&seq, deltaPayload))
if boundary == "delta" {
return b.String()
}
textDonePayload := map[string]any{"type": doneType, "item_id": spec.itemID, "output_index": spec.outputIndex, "content_index": 0, "text": spec.value}
if spec.kind == "message" {
textDonePayload["logprobs"] = []any{}
}
b.WriteString(responsesRawFrame(&seq, textDonePayload))
if boundary == "value_done" {
return b.String()
}
b.WriteString(responsesRawFrame(&seq, map[string]any{"type": "response.content_part.done", "item_id": spec.itemID, "output_index": spec.outputIndex, "content_index": 0, "part": donePart}))
if boundary == "content_done" {
return b.String()
}
completedItem := map[string]any{"id": spec.itemID, "type": spec.kind, "status": "completed", "content": []any{donePart}}
if spec.kind == "message" {
completedItem["role"] = "assistant"
}
b.WriteString(responsesRawFrame(&seq, map[string]any{"type": "response.output_item.done", "output_index": spec.outputIndex, "item": completedItem}))
return b.String()
}
func responsesEventItemID(payload map[string]any) string {
if id, ok := payload["item_id"].(string); ok {
return id
}
if item, ok := payload["item"].(map[string]any); ok {
if id, ok := item["id"].(string); ok {
return id
}
}
return ""
}
func countResponsesItemEvents(payloads []map[string]any, itemID string) map[string]int {
counts := make(map[string]int)
for _, payload := range payloads {
if responsesEventItemID(payload) != itemID {
continue
}
kind, _ := payload["type"].(string)
counts[kind]++
}
return counts
}
func responsesCompletedOutputItem(t *testing.T, payloads []map[string]any, itemID string) map[string]any {
t.Helper()
completed := payloads[len(payloads)-1]
response, _ := completed["response"].(map[string]any)
output, _ := response["output"].([]any)
for _, raw := range output {
item, _ := raw.(map[string]any)
if id, _ := item["id"].(string); id == itemID {
return item
}
}
t.Fatalf("terminal output is missing item %q: %#v", itemID, output)
return nil
}
// TestOpenAIResponsesPoolRawLifecycleContinuation feeds a raw provider prefix
// that opens one message, reasoning, or function item and stops at each distinct
// done boundary, then switches to a normalized continuation of a second item and
// commits the terminal. It proves the raw item keeps its provider identity,
// metadata, and output index, that the continuation item never replaces it, and
// that every remaining lifecycle transition is synthesized exactly once while an
// already-released transition is never repeated.
func TestOpenAIResponsesPoolRawLifecycleContinuation(t *testing.T) {
for _, tc := range []struct {
name string
kind string
boundary string
valueDone string // the item-value done event type for the raw item
}{
{name: "message after delta", kind: "message", boundary: "delta", valueDone: "response.output_text.done"},
{name: "message after text done", kind: "message", boundary: "value_done", valueDone: "response.output_text.done"},
{name: "message after content done", kind: "message", boundary: "content_done", valueDone: "response.output_text.done"},
{name: "message after item done", kind: "message", boundary: "item_done", valueDone: "response.output_text.done"},
{name: "reasoning after delta", kind: "reasoning", boundary: "delta", valueDone: "response.reasoning_text.done"},
{name: "reasoning after text done", kind: "reasoning", boundary: "value_done", valueDone: "response.reasoning_text.done"},
{name: "reasoning after content done", kind: "reasoning", boundary: "content_done", valueDone: "response.reasoning_text.done"},
{name: "reasoning after item done", kind: "reasoning", boundary: "item_done", valueDone: "response.reasoning_text.done"},
{name: "function after delta", kind: "function", boundary: "delta", valueDone: "response.function_call_arguments.done"},
{name: "function after arguments done", kind: "function", boundary: "value_done", valueDone: "response.function_call_arguments.done"},
{name: "function after item done", kind: "function", boundary: "item_done", valueDone: "response.function_call_arguments.done"},
} {
tc := tc
for _, outputIndex := range []int{0, 3} {
outputIndex := outputIndex
t.Run(fmt.Sprintf("%s provider index %d", tc.name, outputIndex), func(t *testing.T) {
w := newRecordingResponseWriter()
selector := newOpenAIStreamGateCodecSelector(openAIStreamGateCodecTunnel)
sink := newOpenAIResponsesPoolReleaseSink(w, &openAIResponsesResultHolder{}, selector)
sink.bindAttempt(true, edgeservice.RunDispatch{Target: "served-a"})
usage := &openAIStreamGateUsageHolder{}
usage.set(usageObservation{inputTokens: 5, outputTokens: 9})
sink.setUsageHolder(usage)
start, err := streamgate.NewResponseStart(streamGateChannelDefault, http.StatusOK, map[string]string{"Content-Type": "text/event-stream"}, time.Now())
if err != nil {
t.Fatalf("NewResponseStart: %v", err)
}
if _, err := sink.CommitResponseStart(context.Background(), start); err != nil {
t.Fatalf("CommitResponseStart: %v", err)
}
spec := rawResponsesItemSpec{kind: tc.kind, responseID: "resp-raw", itemID: "prov-" + tc.kind, outputIndex: outputIndex, value: "raw " + tc.kind}
if tc.kind == "function" {
spec.callID = "call-provider"
spec.name = "provider_fn"
spec.value = `{"raw":true}`
}
sink.codec.pushRelease([]byte(buildRawResponsesItemPrefix(spec, tc.boundary)))
if _, err := sink.Release(context.Background(), mustReleaseText(t, "raw wire ignored")); err != nil {
t.Fatalf("raw prefix release: %v", err)
}
// Normalized text continues an open provider message, but must open a
// fresh message after raw reasoning or function evidence. A provider
// message already closed by its raw lifecycle instead continues with a
// different synthetic item.
sink.bindAttempt(false, edgeservice.RunDispatch{Target: "served-b"})
selector.set(openAIStreamGateCodecNormalized)
continuation := mustReleaseText(t, "continued text")
continuationID := "msg-streamgate-1"
continuesRawMessage := tc.kind == "message" && tc.boundary == "delta"
if continuesRawMessage {
continuationID = spec.itemID
} else if tc.kind == "message" {
continuation = mustReleaseReasoning(t, "continued thought")
continuationID = "reasoning-streamgate-1"
}
if _, err := sink.Release(context.Background(), continuation); err != nil {
t.Fatalf("continuation release: %v", err)
}
terminal, err := streamgate.NewSuccessTerminalResult(streamGateChannelDefault, time.Now())
if err != nil {
t.Fatalf("NewSuccessTerminalResult: %v", err)
}
if _, err := sink.CommitTerminal(context.Background(), terminal); err != nil {
t.Fatalf("CommitTerminal: %v", err)
}
payloads := assertResponsesSSELifecycle(t, w.body.String(), "response.completed")
// The raw provider item keeps its identity and appears in terminal
// output. A provider message is continued in place; other raw kinds
// coexist with a newly opened synthetic message.
rawItem := responsesCompletedOutputItem(t, payloads, spec.itemID)
if rawItem["type"] != mappedItemType(tc.kind) {
t.Fatalf("raw item type = %#v, want %q", rawItem["type"], mappedItemType(tc.kind))
}
if !continuesRawMessage {
if spec.itemID == continuationID {
t.Fatalf("raw item %q collided with continuation item %q", spec.itemID, continuationID)
}
responsesCompletedOutputItem(t, payloads, continuationID) // continuation coexists
continuationAdded := false
for _, payload := range payloads {
if payload["type"] != "response.output_item.added" || responsesEventItemID(payload) != continuationID {
continue
}
if got := int(payload["output_index"].(float64)); got != outputIndex+1 {
t.Fatalf("continuation output_index = %d, want %d", got, outputIndex+1)
}
continuationAdded = true
}
if !continuationAdded {
t.Fatalf("continuation item %q was not opened", continuationID)
}
}
if tc.kind == "function" {
if rawItem["call_id"] != spec.callID || rawItem["name"] != spec.name {
t.Fatalf("raw function metadata = (%#v,%#v), want (%q,%q)", rawItem["call_id"], rawItem["name"], spec.callID, spec.name)
}
if rawItem["arguments"] != spec.value {
t.Fatalf("raw function arguments = %#v, want %q", rawItem["arguments"], spec.value)
}
} else {
content, _ := rawItem["content"].([]any)
if len(content) != 1 {
t.Fatalf("raw content item invalid: %#v", rawItem)
}
part, _ := content[0].(map[string]any)
want := spec.value
if continuesRawMessage {
want += "continued text"
}
if part["text"] != want {
t.Fatalf("raw content text = %#v, want %q", part["text"], want)
}
}
// Every lifecycle-closing transition for the raw item is present exactly
// once, whether it came from the raw prefix or was synthesized.
counts := countResponsesItemEvents(payloads, spec.itemID)
if counts[tc.valueDone] != 1 {
t.Fatalf("%s count for %q = %d, want exactly 1; payloads=%#v", tc.valueDone, spec.itemID, counts[tc.valueDone], payloads)
}
if counts["response.output_item.done"] != 1 {
t.Fatalf("output_item.done count for %q = %d, want exactly 1", spec.itemID, counts["response.output_item.done"])
}
if tc.kind != "function" && counts["response.content_part.done"] != 1 {
t.Fatalf("content_part.done count for %q = %d, want exactly 1", spec.itemID, counts["response.content_part.done"])
}
})
}
}
}
func mappedItemType(kind string) string {
if kind == "function" {
return "function_call"
}
return kind
}
// TestStreamGateResponsesPoolStreamingTunnelReleasesIncrementally isolates the
// no-recovery path. It proves the caller-visible initial tunnel is not routed
// through a buffered delegate while the terminal frame is still withheld.
func TestStreamGateResponsesPoolStreamingTunnelReleasesIncrementally(t *testing.T) {
frames := make(chan *iop.ProviderTunnelFrame)
service := newScriptedPoolRunService(scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathTunnel), runID: "resp-incremental", provider: "prov-a", target: "served-a", frames: frames,
})
srv, requestCtx, poolReq := buildStreamGateResponsesPoolFixture(t, service, 3)
result, err := service.SubmitProviderPool(context.Background(), poolReq)
if err != nil {
t.Fatalf("initial SubmitProviderPool: %v", err)
}
w := newNotifyingResponseWriter()
done := make(chan struct{})
go func() {
srv.runOpenAIResponsesPoolStreamGate(w, requestCtx, poolReq, result.Tunnel)
close(done)
}()
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_RESPONSE_START, StatusCode: http.StatusOK, Headers: map[string]string{"Content-Type": "text/event-stream"}}
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: []byte(responsesStreamPrefix("resp-incremental", "msg-incremental", "first safe frame"))}
select {
case <-w.writes:
case <-time.After(5 * time.Second):
t.Fatal("initial streaming Responses frame was not written before terminal")
}
select {
case <-done:
t.Fatal("Responses runtime completed before the terminal frame was provided")
default:
}
body := w.body.String()
if got := responsesSSEOutputText(parseCompleteResponsesSSE(t, body)); got != "first safe frame" {
t.Fatalf("pre-terminal Responses SSE output = %q, body=%q", got, body)
}
frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: []byte(responsesStreamTerminal("resp-incremental", "msg-incremental", "first safe frame", 5))}
close(frames)
select {
case <-done:
case <-time.After(5 * time.Second):
t.Fatal("Responses runtime did not finish after terminal frame")
}
if got := responsesSSEOutputText(assertResponsesSSELifecycle(t, w.body.String(), "response.completed")); got != "first safe frame" {
t.Fatalf("completed Responses SSE output = %q", got)
}
}
func TestStreamGateResponsesPoolInitialErrorPreservesProviderResponse(t *testing.T) {
providerBody := `{"error":{"message":"provider capacity exhausted","type":"rate_limit_error"}}`
service := newScriptedPoolRunService(scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathTunnel), runID: "resp-provider-error", provider: "prov-a", target: "served-a",
frames: bufferedTunnelFrames(
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_RESPONSE_START, StatusCode: http.StatusTooManyRequests, Headers: map[string]string{"Content-Type": "application/json", "X-Provider-Trace": "trace-123"}},
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: []byte(providerBody)},
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_END, End: true},
),
})
srv, requestCtx, poolReq := buildStreamGateResponsesPoolFixture(t, service, 3)
result, err := service.SubmitProviderPool(context.Background(), poolReq)
if err != nil {
t.Fatalf("initial SubmitProviderPool: %v", err)
}
w := newRecordingResponseWriter()
srv.runOpenAIResponsesPoolStreamGate(w, requestCtx, poolReq, result.Tunnel)
if got := w.code; got != http.StatusTooManyRequests {
t.Fatalf("provider status = %d, want %d", got, http.StatusTooManyRequests)
}
if got := w.Header().Get("Content-Type"); got != "application/json" {
t.Fatalf("provider Content-Type = %q, want application/json", got)
}
if got := w.Header().Get("X-Provider-Trace"); got != "trace-123" {
t.Fatalf("provider trace header = %q", got)
}
if got := w.body.String(); got != providerBody {
t.Fatalf("provider body = %q, want byte-for-byte %q", got, providerBody)
}
if got := w.headerCallCount(); got != 1 {
t.Fatalf("WriteHeader call count = %d, want 1", got)
}
if strings.Contains(w.body.String(), "data:") {
t.Fatalf("pre-commit provider error was rewritten as SSE: %q", w.body.String())
}
}
func TestStreamGateResponsesPoolRecoveryRejectionAfterStreamOpenTerminatesSSE(t *testing.T) {
service := newScriptedPoolRunService(
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathTunnel), runID: "resp-open-prefix", provider: "prov-a", target: "served-a",
frames: responsesTunnelFrames(responsesStreamPrefix("resp-open", "msg-open", "safe prefix ") +
"data: {\"type\":\"response.output_text.delta\",\"item_id\":\"msg-open\",\"output_index\":0,\"content_index\":0,\"delta\":\"discarded repeat tail\",\"sequence_number\":5}\n\n"),
},
scriptedPoolAttempt{err: edgeservice.ErrProviderPoolCandidateRejected},
)
srv, requestCtx, poolReq := buildStreamGateResponsesPoolFixture(t, service, 3)
result, err := service.SubmitProviderPool(context.Background(), poolReq)
if err != nil {
t.Fatalf("initial SubmitProviderPool: %v", err)
}
dc := newOpenAIResponsesPoolTunnelDispatchContext(requestCtx, poolReq)
snapRef, err := requestCtx.ingress.recoveryRef()
if err != nil {
t.Fatalf("recoveryRef: %v", err)
}
violation := newInjectedContinuationViolationFilter(t, "test.responses.rejection", 1, snapRef.SnapshotRef(), len("safe prefix "))
reg, err := streamgate.NewFilterRegistration(violation, streamGateNoopCapability, true, streamgate.FilterEnforcementBlocking, streamGateFilterTimeout, injectedViolationPriority)
if err != nil {
t.Fatalf("NewFilterRegistration: %v", err)
}
w := newRecordingResponseWriter()
holder := &openAIResponsesResultHolder{}
sink := newOpenAIResponsesPoolReleaseSink(w, holder, newOpenAIStreamGateCodecSelector(openAIStreamGateCodecTunnel))
runtime, _, err := srv.buildOpenAIResponsesStreamGateRuntimeFromAttempt(
dc, openAIAttemptTransport{path: openAIAdmissionTunnel, tunnel: result.Tunnel}, result.Tunnel.Dispatch(), result.Tunnel.Close,
sink, streamGateTestRegistry(t, reg),
)
if err != nil {
t.Fatalf("buildOpenAIResponsesStreamGateRuntimeFromAttempt: %v", err)
}
if err := runtime.Run(context.Background()); err != nil {
t.Fatalf("runtime.Run: %v", err)
}
body := w.body.String()
payloads := assertResponsesSSELifecycle(t, body, "error")
if got := responsesSSEOutputText(payloads); got != "safe prefix " {
t.Fatalf("released prefix = %q, want safe prefix only; body=%q", got, body)
}
if !strings.Contains(body, openAIStreamGateCandidateRejectedMessage) {
t.Fatalf("candidate rejection error missing from SSE: %q", body)
}
if w.code != http.StatusOK || w.Header().Get("Content-Type") != "text/event-stream" {
t.Fatalf("post-open rejection status/header = %d/%q, want 200/text-event-stream", w.code, w.Header().Get("Content-Type"))
}
if got := w.headerCallCount(); got != 1 {
t.Fatalf("WriteHeader call count = %d, want 1", got)
}
if strings.Contains(body, `{"error":{`) {
t.Fatalf("post-open rejection appended HTTP JSON: %q", body)
}
}
// --- S16: provider-pool + tool validation (one owner across a path switch) ---
const streamGatePoolToolBody = `{"model":"client-model","messages":[{"role":"user","content":"status"}],` +
`"tools":[{"type":"function","function":{"name":"run_commands","parameters":{"type":"object","properties":{"commands":{"type":"array","items":{"type":"string"}}}}}}]}`
// TestStreamGatePoolToolValidationReAdmitsThroughPool proves the tool-validation
// terminal gate and the provider-pool re-admission are the same single cycle:
// an invalid tool call on a normalized pool candidate re-enters
// SubmitProviderPool once, and the passing candidate's structured output is the
// only thing rendered.
func TestStreamGatePoolToolValidationReAdmitsThroughPool(t *testing.T) {
service := newScriptedPoolRunService(
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathNormalized), runID: "pool-tool-1",
provider: "prov-a", target: "served-a",
runEvents: bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "<tool_call><function=delete_everything><parameter=confirm>true</parameter></function></tool_call>"},
&iop.RunEvent{Type: "complete"},
),
},
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathNormalized), runID: "pool-tool-2",
provider: "prov-b", target: "served-b",
runEvents: bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "<tool_call><function=run_commands><parameter=commands>[\"ls\"]\n</parameter></function></tool_call>"},
&iop.RunEvent{Type: "complete"},
),
},
)
srv, dc := buildStreamGatePoolChatFixture(t, service, []byte(streamGatePoolToolBody), 1)
if !dc.validation.enabled {
t.Fatal("fixture must carry a tool-validation contract")
}
w, sink, runErr := runPoolChatStreamGate(t, srv, dc)
if runErr != nil {
t.Fatalf("rt.Run: %v", runErr)
}
if got := service.poolSubmits(); got != 2 {
t.Fatalf("SubmitProviderPool calls: got %d, want 2 (tool-validation recovery must re-enter the pool)", got)
}
if got := sink.(*openAICompositeReleaseSink).resolvedCodec(); got != openAIStreamGateCodecNormalized {
t.Fatalf("resolved codec: got %q, want normalized", got)
}
var resp chatCompletionResponse
if err := json.Unmarshal([]byte(w.body.String()), &resp); err != nil {
t.Fatalf("decode response: %v body=%s", err, w.body.String())
}
if len(resp.Choices) != 1 || resp.Choices[0].FinishReason != "tool_calls" {
t.Fatalf("expected a structured tool_calls completion, got %s", w.body.String())
}
if len(resp.Choices[0].Message.ToolCalls) != 1 {
t.Fatalf("tool calls: got %d, want 1: %s", len(resp.Choices[0].Message.ToolCalls), w.body.String())
}
if strings.Contains(w.body.String(), "delete_everything") {
t.Fatalf("the rejected attempt leaked into the response: %s", w.body.String())
}
}
// TestStreamGatePoolToolValidationDropsOnTunnelCandidate verifies the
// tool-validation contract is re-resolved per attempt: after a recovery selects
// a raw passthrough candidate the normalized contract no longer applies, and
// the response is relayed with the tunnel framing instead of being validated
// again.
func TestStreamGatePoolToolValidationDropsOnTunnelCandidate(t *testing.T) {
service := newScriptedPoolRunService(
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathNormalized), runID: "pool-tool-1",
provider: "prov-a", target: "served-a",
runEvents: bufferedRunEvents(
&iop.RunEvent{Type: "delta", Delta: "<tool_call><function=delete_everything><parameter=confirm>true</parameter></function></tool_call>"},
&iop.RunEvent{Type: "complete"},
),
},
scriptedPoolAttempt{
path: string(edgeservice.ProviderPoolPathTunnel), runID: "pool-tool-2",
provider: "prov-b", target: "served-b",
frames: bufferedTunnelFrames(
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_RESPONSE_START, StatusCode: 200, Headers: map[string]string{"Content-Type": "application/json"}},
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: []byte(`{"id":"provider-native","object":"chat.completion"}`)},
&iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_END, End: true},
),
},
)
srv, dc := buildStreamGatePoolChatFixture(t, service, []byte(streamGatePoolToolBody), 1)
w, sink, runErr := runPoolChatStreamGate(t, srv, dc)
if runErr != nil {
t.Fatalf("rt.Run: %v", runErr)
}
if got := service.poolSubmits(); got != 2 {
t.Fatalf("SubmitProviderPool calls: got %d, want 2", got)
}
if got := sink.(*openAICompositeReleaseSink).resolvedCodec(); got != openAIStreamGateCodecTunnel {
t.Fatalf("resolved codec: got %q, want tunnel", got)
}
body := w.body.String()
if !strings.Contains(body, `"id":"provider-native"`) {
t.Fatalf("raw passthrough bytes missing: %q", body)
}
if strings.Contains(body, "delete_everything") {
t.Fatalf("the rejected normalized attempt leaked into the response: %q", body)
}
if got := w.headerCallCount(); got != 1 {
t.Fatalf("WriteHeader call count: got %d, want 1", got)
}
}