From 69cb1351dd2cac6d4ff2be60cec90999cbed52d0 Mon Sep 17 00:00:00 2001 From: tbphp Date: Wed, 9 Sep 2026 19:14:47 +0800 Subject: [PATCH 1/3] =?UTF-8?q?feat(responses):=20=E5=AE=9E=E7=8E=B0?= =?UTF-8?q?=E5=93=8D=E5=BA=94=E7=8A=B6=E6=80=81=E7=BB=AD=E6=8E=A5=E8=B7=AF?= =?UTF-8?q?=E7=94=B1?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.md | 4 +- README_CN.md | 4 +- README_JP.md | 4 +- internal/app/runtime_checkpoint.go | 28 +- internal/app/runtime_checkpoint_test.go | 34 +- internal/app/scheduling_checkpoint_test.go | 4 +- internal/container/container.go | 4 +- internal/control/group_create.go | 7 + internal/control/group_settings_test.go | 32 ++ internal/dialect/anthropic.go | 2 +- internal/dialect/dialect.go | 1 + internal/dialect/openai.go | 2 +- internal/dialect/openai_embeddings.go | 2 +- internal/dialect/openai_images.go | 2 +- internal/dialect/openai_responses.go | 6 +- internal/dialect/openai_responses_test.go | 33 ++ internal/dialect/request_fields.go | 11 +- internal/dialect/request_fields_test.go | 2 +- internal/dialect/rerank.go | 2 +- internal/gateway/execution_forward.go | 13 +- internal/gateway/forward.go | 2 + internal/gateway/handler.go | 76 +++-- internal/gateway/reason.go | 1 + internal/gateway/reason_test.go | 1 + internal/gateway/responses_continuation.go | 46 +++ .../gateway/responses_continuation_test.go | 307 ++++++++++++++++++ .../response_continuation_test.go | 33 ++ internal/parameteroverride/rules.go | 31 ++ internal/state/response_bindings.go | 158 +++++++++ internal/state/response_bindings_test.go | 128 ++++++++ 30 files changed, 930 insertions(+), 50 deletions(-) create mode 100644 internal/gateway/responses_continuation.go create mode 100644 internal/gateway/responses_continuation_test.go create mode 100644 internal/parameteroverride/response_continuation_test.go create mode 100644 internal/state/response_bindings.go create mode 100644 internal/state/response_bindings_test.go diff --git a/README.md b/README.md index e199d9e3f..00f14a977 100644 --- a/README.md +++ b/README.md @@ -239,7 +239,9 @@ Environment proxies apply only when no proxy is specified on the credential, gro - 2.0 is designed for a **single application instance**. Instances do not share state, so horizontal scaling is not supported. - Usage and cost are **estimates** derived from upstream responses. They support operational analysis and capacity planning, and do not equal a provider invoice or a financial reconciliation. - Subscription channels depend on upstream OAuth and compatibility protocols and may change as upstreams change. Only connect accounts you are entitled to use, and follow each provider's terms. -- In OpenAI Responses, stateful requests relying on `previous_response_id`, `conversation`, or an existing resource ID are only reliable with a single credential, or with an upstream that shares resources across credentials. +- Responses continuation with `previous_response_id` currently supports the `openai` and `gpt_load` channels. Response ownership is isolated by AccessKey and pins the original credential when current routing permits, independently of soft affinity. Unknown IDs, including IDs created before upgrading or outside this gateway, are rejected. Group parameter overrides cannot change this field. +- Response bindings stay in memory for up to 30 days, with limits of 100,000 entries and 16 MiB of ID text; older entries are evicted when capacity is reached. A successful checkpoint during normal shutdown allows restoration from the same data directory. Crash recovery and continued upstream state availability are not guaranteed. +- `conversation` and other existing resource IDs are outside this ownership routing scope and still depend on a single credential or upstream resource sharing across credentials. ## Moving from 1.x diff --git a/README_CN.md b/README_CN.md index dad9a314f..acadf4e7b 100644 --- a/README_CN.md +++ b/README_CN.md @@ -239,7 +239,9 @@ Windows 普通用户可改为下载 `gpt-load-windows-setup.exe`。双击并确 - 2.0 按**单应用实例**设计,多个实例之间不共享状态,不支持直接横向扩容。 - 用量与成本是基于上游返回数据的**估算**,用于运行分析和资源评估,不等同于服务商账单或财务对账结果。 - 订阅渠道依赖上游 OAuth 与兼容协议,可能随上游变化调整。请只接入自己有权使用的账号,并遵守对应服务商条款。 -- OpenAI Responses 中依赖 `previous_response_id`、`conversation` 或既有资源 ID 的有状态请求,只有在单凭据或上游跨凭据共享资源时才可靠。 +- Responses 的 `previous_response_id` 续接目前支持 `openai`、`gpt_load` 渠道:按 AccessKey 隔离响应归属,并在当前路由允许时固定原凭据,不受软亲和开关影响。未知 ID(包括升级前或网关外创建的 ID)直接拒绝;Group 参数覆盖不能改写该字段。 +- 响应归属保存在内存中,默认保留 30 天,最多 100,000 条,ID 文本合计最多 16 MiB,达到容量时淘汰旧记录。正常停机成功保存 checkpoint 后可在同一数据目录恢复;不保证崩溃恢复或上游历史仍有效。 +- `conversation` 与其他既有资源 ID 不在上述归属路由范围内,仍依赖单凭据或上游跨凭据共享资源。 ## 从 1.x 切换 diff --git a/README_JP.md b/README_JP.md index 942179b77..384d65e2f 100644 --- a/README_JP.md +++ b/README_JP.md @@ -239,7 +239,9 @@ Windows の一般ユーザーは代わりに `gpt-load-windows-setup.exe` を利 - 2.0 は**単一アプリケーションインスタンス**を前提に設計されています。インスタンス間で状態を共有しないため、そのままの水平スケールには対応していません。 - 使用量とコストはアップストリームの応答に基づく**概算**です。運用分析やリソース評価には使えますが、プロバイダーの請求書や会計上の照合結果とは一致しません。 - サブスクリプションチャネルはアップストリームの OAuth と互換プロトコルに依存し、アップストリームの変更に伴って調整が必要になる場合があります。利用権限のあるアカウントのみを接続し、各プロバイダーの規約に従ってください。 -- OpenAI Responses で `previous_response_id`、`conversation`、既存のリソース ID に依存するステートフルなリクエストは、単一の認証情報を使う場合、またはアップストリームが認証情報間でリソースを共有している場合にのみ確実に動作します。 +- Responses の `previous_response_id` による継続は、現在 `openai`、`gpt_load` チャネルに対応しています。応答の帰属を AccessKey ごとに分離し、現在のルーティングで許可される場合に元の認証情報へ固定します。ソフトアフィニティ設定には依存しません。アップグレード前やゲートウェイ外で作成されたものを含め、不明な ID は拒否されます。Group のパラメータ上書きでこのフィールドを変更することはできません。 +- 応答の帰属はメモリに最大 30 日間保持され、上限は 100,000 件および ID テキスト合計 16 MiB です。容量に達すると古い記録を削除します。通常終了時に checkpoint の保存が成功すれば、同じデータディレクトリから復元できます。クラッシュからの復元や、アップストリームの履歴が引き続き有効であることは保証しません。 +- `conversation` とその他の既存リソース ID はこの帰属ルーティングの対象外であり、単一の認証情報またはアップストリームでの認証情報間のリソース共有が引き続き必要です。 ## 1.x からの移行 diff --git a/internal/app/runtime_checkpoint.go b/internal/app/runtime_checkpoint.go index 02bf60d57..9f73b84f9 100644 --- a/internal/app/runtime_checkpoint.go +++ b/internal/app/runtime_checkpoint.go @@ -25,28 +25,32 @@ type runtimeStateCheckpointDocument struct { Credentials []state.CredentialRuntimeCheckpoint `json:"credentials,omitempty"` Stats []health.StatsRuntimeCheckpoint `json:"stats,omitempty"` Scheduling *state.SchedulingCheckpoint `json:"scheduling,omitempty"` + Responses []state.ResponseBinding `json:"responses,omitempty"` } // FileRuntimeStateCheckpoint stores the small, disposable runtime checkpoint // in DATA_DIR. The startup path consumes the file before parsing it so a // malformed or partially written file cannot be retried on every restart. type FileRuntimeStateCheckpoint struct { - path string - registry *state.CredentialRegistry - stats *health.StatsStore - removeFile func(string) error + path string + registry *state.CredentialRegistry + stats *health.StatsStore + responseBindings *state.ResponseBindings + removeFile func(string) error } func NewFileRuntimeStateCheckpoint( dataDir string, registry *state.CredentialRegistry, stats *health.StatsStore, + responseBindings *state.ResponseBindings, ) *FileRuntimeStateCheckpoint { return &FileRuntimeStateCheckpoint{ - path: filepath.Join(dataDir, runtimeStateCheckpointFileName), - registry: registry, - stats: stats, - removeFile: os.Remove, + path: filepath.Join(dataDir, runtimeStateCheckpointFileName), + registry: registry, + stats: stats, + responseBindings: responseBindings, + removeFile: os.Remove, } } @@ -80,6 +84,11 @@ func (checkpoint *FileRuntimeStateCheckpoint) Restore(ctx context.Context) error if checkpoint.stats != nil { checkpoint.stats.RestoreRuntimeCheckpoint(document.Stats) } + if checkpoint.responseBindings != nil { + if err := checkpoint.responseBindings.RestoreCheckpoint(document.Responses); err != nil { + return err + } + } return nil } @@ -99,6 +108,9 @@ func (checkpoint *FileRuntimeStateCheckpoint) Save(ctx context.Context) error { if checkpoint.stats != nil { document.Stats = checkpoint.stats.CaptureRuntimeCheckpoint() } + if checkpoint.responseBindings != nil { + document.Responses = checkpoint.responseBindings.CaptureCheckpoint() + } payload, err := json.Marshal(document) if err != nil { return err diff --git a/internal/app/runtime_checkpoint_test.go b/internal/app/runtime_checkpoint_test.go index 87f904375..4c0d6c6e3 100644 --- a/internal/app/runtime_checkpoint_test.go +++ b/internal/app/runtime_checkpoint_test.go @@ -110,7 +110,7 @@ func TestAppLogsCheckpointRestoreFailureOnce(t *testing.T) { DB: db, StartupBootstrap: startupBootstrapFunc(noopStartupBootstrap), RuntimeState: runtimeStateLoaderFunc(func(context.Context) error { return nil }), - RuntimeCheckpoint: NewFileRuntimeStateCheckpoint(dataDir, nil, nil), + RuntimeCheckpoint: NewFileRuntimeStateCheckpoint(dataDir, nil, nil, nil), ControlRuntime: newControlRuntimeFake(nil, false), RequestLogs: newRequestLogRuntimeFake(nil, nil), }) @@ -199,7 +199,7 @@ func TestFileRuntimeStateCheckpointRestoresAndConsumesFile(t *testing.T) { stats := health.NewStatsStore() stats.RecordFailure(1, health.FailureCategoryUpstreamHostError, 503, time.Date(2026, 8, 7, 11, 59, 0, 0, time.UTC)) - checkpoint := NewFileRuntimeStateCheckpoint(dataDir, registry, stats) + checkpoint := NewFileRuntimeStateCheckpoint(dataDir, registry, stats, nil) if err := checkpoint.Save(context.Background()); err != nil { t.Fatalf("Save() error = %v", err) } @@ -218,7 +218,7 @@ func TestFileRuntimeStateCheckpointRestoresAndConsumesFile(t *testing.T) { t.Fatalf("replace loaded registry: %v", err) } loadedStats := health.NewStatsStore() - loader := NewFileRuntimeStateCheckpoint(dataDir, loadedRegistry, loadedStats) + loader := NewFileRuntimeStateCheckpoint(dataDir, loadedRegistry, loadedStats, nil) if err := loader.Restore(context.Background()); err != nil { t.Fatalf("Restore() error = %v", err) } @@ -237,6 +237,30 @@ func TestFileRuntimeStateCheckpointRestoresAndConsumesFile(t *testing.T) { } } +func TestFileRuntimeStateCheckpointRestoresResponseOwnership(t *testing.T) { + dir := t.TempDir() + original := state.NewResponseBindings() + if !original.Record(7, "stored-response", state.CredentialRef{ID: 2, GroupID: 3, IdentityGeneration: 4}) { + t.Fatal("record failed") + } + want, _ := original.Lookup(7, "stored-response") + checkpoint := NewFileRuntimeStateCheckpoint(dir, nil, nil, original) + if err := checkpoint.Save(context.Background()); err != nil { + t.Fatal(err) + } + restored := state.NewResponseBindings() + loader := NewFileRuntimeStateCheckpoint(dir, nil, nil, restored) + if err := loader.Restore(context.Background()); err != nil { + t.Fatal(err) + } + got, ok := restored.Lookup(7, "stored-response") + if !ok || got.AccessKeyID != want.AccessKeyID || got.CredentialID != want.CredentialID || + got.GroupID != want.GroupID || got.IdentityGeneration != want.IdentityGeneration || + !got.ExpiresAt.Equal(want.ExpiresAt) { + t.Fatalf("restored binding = %#v, %t; want %#v", got, ok, want) + } +} + func TestFileRuntimeStateCheckpointReturnsErrorWhenDeleteFails(t *testing.T) { dataDir := t.TempDir() path := filepath.Join(dataDir, runtimeStateCheckpointFileName) @@ -257,7 +281,7 @@ func TestFileRuntimeStateCheckpointReturnsErrorWhenDeleteFails(t *testing.T) { }}); err != nil { t.Fatalf("replace registry: %v", err) } - checkpoint := NewFileRuntimeStateCheckpoint(dataDir, registry, health.NewStatsStore()) + checkpoint := NewFileRuntimeStateCheckpoint(dataDir, registry, health.NewStatsStore(), nil) // The normal file implementation removes the file successfully. This test // documents that a failed removal must prevent applying stale data through // the injectable filesystem hook used by the implementation. @@ -283,7 +307,7 @@ func TestFileRuntimeStateCheckpointConsumesMalformedFileAndReturnsError(t *testi }}); err != nil { t.Fatalf("replace registry: %v", err) } - checkpoint := NewFileRuntimeStateCheckpoint(dataDir, registry, health.NewStatsStore()) + checkpoint := NewFileRuntimeStateCheckpoint(dataDir, registry, health.NewStatsStore(), nil) if err := checkpoint.Restore(context.Background()); err == nil { t.Fatal("Restore() error = nil, want malformed checkpoint error") } diff --git a/internal/app/scheduling_checkpoint_test.go b/internal/app/scheduling_checkpoint_test.go index 65c368532..91e42d7d9 100644 --- a/internal/app/scheduling_checkpoint_test.go +++ b/internal/app/scheduling_checkpoint_test.go @@ -31,11 +31,11 @@ func TestRuntimeCheckpointRestoresOnlyMatchingSchedulingIdentities(t *testing.T) d.Members[2].Progress = d.Watermark d.Members[2].LastSelected = 8 }) - if err := NewFileRuntimeStateCheckpoint(dir, original, nil).Save(context.Background()); err != nil { + if err := NewFileRuntimeStateCheckpoint(dir, original, nil, nil).Save(context.Background()); err != nil { t.Fatal(err) } loaded := makeRegistry(map[uint]uint64{1: 1, 2: 2, 3: 1}) - if err := NewFileRuntimeStateCheckpoint(dir, loaded, nil).Restore(context.Background()); err != nil { + if err := NewFileRuntimeStateCheckpoint(dir, loaded, nil, nil).Restore(context.Background()); err != nil { t.Fatal(err) } loaded.SchedulingState().WithLock(func(d *state.SchedulingLedger) { diff --git a/internal/container/container.go b/internal/container/container.go index bb11f66ae..4e875131b 100644 --- a/internal/container/container.go +++ b/internal/container/container.go @@ -65,6 +65,7 @@ func BuildContainer() (*dig.Container, error) { app.NewEngineWithLifecycle, webui.NewServer, state.NewCredentialRegistry, + state.NewResponseBindings, accessquota.NewRuntime, channel.CompileRegistry, control.NewPriceRuntime, @@ -118,8 +119,9 @@ func BuildContainer() (*dig.Container, error) { cfg *config.Config, registry *state.CredentialRegistry, stats *health.StatsStore, + responseBindings *state.ResponseBindings, ) app.RuntimeStateCheckpoint { - return app.NewFileRuntimeStateCheckpoint(cfg.DataDir, registry, stats) + return app.NewFileRuntimeStateCheckpoint(cfg.DataDir, registry, stats, responseBindings) }, control.NewRuntime, func(runtime *control.Runtime) app.ControlRuntime { return runtime }, diff --git a/internal/control/group_create.go b/internal/control/group_create.go index 815ea8133..1d2e39de3 100644 --- a/internal/control/group_create.go +++ b/internal/control/group_create.go @@ -13,6 +13,7 @@ import ( "gpt-load/internal/channel" "gpt-load/internal/outboundproxy" + "gpt-load/internal/parameteroverride" "gpt-load/internal/platform/config" app_errors "gpt-load/internal/platform/errors" "gpt-load/internal/platform/utils" @@ -307,6 +308,12 @@ func normalizeGroupSettings(settings config.Settings) (config.Settings, models.J if settings == nil { settings = make(config.Settings) } + if value, exists := settings[state.SettingParameterOverrides]; exists { + rules, err := parameteroverride.Compile(value) + if err != nil || rules.ValidateResponsesContinuation() != nil { + return nil, nil, app_errors.ErrValidation + } + } encoded, err := json.Marshal(settings) if err != nil { return nil, nil, app_errors.ErrValidation diff --git a/internal/control/group_settings_test.go b/internal/control/group_settings_test.go index 160c64c60..8b700d408 100644 --- a/internal/control/group_settings_test.go +++ b/internal/control/group_settings_test.go @@ -119,6 +119,38 @@ func TestUpdateGroupSettingsPublishesParameterOverrides(t *testing.T) { } } +func TestGroupSettingsRejectNewContinuationOverridesButKeepLegacyReadable(t *testing.T) { + fixture := newServiceFixture(t) + group := validControlGroup("legacy-continuation") + group.Overrides = models.JSON(`{"parameter_overrides":[{"set":{"previous_response_id":"old-response"}}]}`) + if err := fixture.db.Create(group).Error; err != nil { + t.Fatal(err) + } + if _, err := fixture.manager.Publish(mustBuildCompileInput(t, fixture.db)); err != nil { + t.Fatalf("legacy rules blocked snapshot loading: %v", err) + } + if _, err := fixture.service.GetGroupSettings(t.Context(), group.ID); err != nil { + t.Fatalf("legacy rules prevented management access: %v", err) + } + before := fixture.manager.Current() + _, err := fixture.service.UpdateGroupSettings(t.Context(), group.ID, GroupSettingsUpdateRequest{ + Overrides: optionalField[config.Settings]{Set: true, Value: config.Settings{ + state.SettingParameterOverrides: []any{map[string]any{ + "set": map[string]any{"previous_response_id": "new-response"}, + }}, + }}, + }) + if !errors.Is(err, app_errors.ErrValidation) || fixture.manager.Current() != before { + t.Fatalf("new continuation override error = %v", err) + } + _, err = fixture.service.UpdateGroupSettings(t.Context(), group.ID, GroupSettingsUpdateRequest{ + Overrides: optionalField[config.Settings]{Set: true, Value: config.Settings{}}, + }) + if err != nil { + t.Fatalf("legacy rule could not be corrected: %v", err) + } +} + func TestUpdateGroupSettingsPublishesOnceAndReturnsNewSettings(t *testing.T) { t.Parallel() fixture := newServiceFixture(t) diff --git a/internal/dialect/anthropic.go b/internal/dialect/anthropic.go index 69b4588ff..25f3d3bf9 100644 --- a/internal/dialect/anthropic.go +++ b/internal/dialect/anthropic.go @@ -42,7 +42,7 @@ func (d *Anthropic) InspectRequest(req *ParsedRequest) (RequestMetadata, error) return RequestMetadata{}, fmt.Errorf("parsed request is required") } - metadata, err := inspectJSONRequestFields(req.Body, true) + metadata, err := inspectJSONRequestFields(req.Body, true, false) if err != nil { return RequestMetadata{}, fmt.Errorf("decode %s request: %w", d.Protocol(), err) } diff --git a/internal/dialect/dialect.go b/internal/dialect/dialect.go index ea9173d9c..d99c00a40 100644 --- a/internal/dialect/dialect.go +++ b/internal/dialect/dialect.go @@ -25,6 +25,7 @@ type RequestMetadata struct { Operation execution.Operation RouteRequirement execution.RouteRequirement ResponsesStorePreference execution.ResponsesStorePreference + PreviousResponseID string ObserveUsage bool PricingMode pricing.Mode UsageDiagnostics usage.Diagnostics diff --git a/internal/dialect/openai.go b/internal/dialect/openai.go index 272924754..ba4a26a25 100644 --- a/internal/dialect/openai.go +++ b/internal/dialect/openai.go @@ -23,7 +23,7 @@ func (d *OpenAI) InspectRequest(req *ParsedRequest) (RequestMetadata, error) { return RequestMetadata{}, fmt.Errorf("parsed request is required") } - metadata, err := inspectJSONRequestFields(req.Body, true) + metadata, err := inspectJSONRequestFields(req.Body, true, false) if err != nil { return RequestMetadata{}, fmt.Errorf("decode %s request: %w", d.Protocol(), err) } diff --git a/internal/dialect/openai_embeddings.go b/internal/dialect/openai_embeddings.go index 22c2ed5b0..7e5a69557 100644 --- a/internal/dialect/openai_embeddings.go +++ b/internal/dialect/openai_embeddings.go @@ -51,7 +51,7 @@ func (d *OpenAIEmbeddings) InspectRequest(request *ParsedRequest) (RequestMetada return RequestMetadata{}, fmt.Errorf("unsupported %s Content-Type %q", d.Protocol(), contentType) } } - metadata, err := inspectJSONRequestFields(request.Body, true) + metadata, err := inspectJSONRequestFields(request.Body, true, false) if err != nil { return RequestMetadata{}, fmt.Errorf("decode %s request: %w", d.Protocol(), err) } diff --git a/internal/dialect/openai_images.go b/internal/dialect/openai_images.go index b312c5fd0..bd4c9c0d6 100644 --- a/internal/dialect/openai_images.go +++ b/internal/dialect/openai_images.go @@ -74,7 +74,7 @@ func (d *OpenAIImages) InspectRequest(request *ParsedRequest) (RequestMetadata, if mediaType != "" && mediaType != "application/json" { return RequestMetadata{}, fmt.Errorf("unsupported %s Content-Type %q", d.Protocol(), contentType) } - metadata, err := inspectJSONRequestFields(request.Body, true) + metadata, err := inspectJSONRequestFields(request.Body, true, false) if err != nil { return RequestMetadata{}, fmt.Errorf("decode %s request: %w", d.Protocol(), err) } diff --git a/internal/dialect/openai_responses.go b/internal/dialect/openai_responses.go index ecef3c887..a3cf2df8e 100644 --- a/internal/dialect/openai_responses.go +++ b/internal/dialect/openai_responses.go @@ -38,7 +38,7 @@ func (d *OpenAIResponses) InspectRequest(req *ParsedRequest) (RequestMetadata, e metadata := RequestMetadata{} if len(req.Body) > 0 { - parsed, err := inspectJSONRequestFields(req.Body, false) + parsed, err := inspectJSONRequestFields(req.Body, false, req.Method == http.MethodPost && req.Path == openAIResponsesPath) if err != nil { return RequestMetadata{}, fmt.Errorf("decode %s request: %w", d.Protocol(), err) } @@ -56,7 +56,9 @@ func (d *OpenAIResponses) InspectRequest(req *ParsedRequest) (RequestMetadata, e metadata.ObserveUsage = req.Method == http.MethodPost && (req.Path == openAIResponsesPath || req.Path == openAIResponsesCompactPath) if len(req.Body) > 0 { - metadata.AffinityPrefix = inspectPromptAffinityPrefix(d.Protocol(), req.Body) + if metadata.PreviousResponseID == "" { + metadata.AffinityPrefix = inspectPromptAffinityPrefix(d.Protocol(), req.Body) + } pricingMode, diagnostics, err := openAIRequestPricing(req.Body) if err != nil { return RequestMetadata{}, fmt.Errorf("inspect %s request pricing: %w", d.Protocol(), err) diff --git a/internal/dialect/openai_responses_test.go b/internal/dialect/openai_responses_test.go index dd3108126..dfef173ce 100644 --- a/internal/dialect/openai_responses_test.go +++ b/internal/dialect/openai_responses_test.go @@ -78,3 +78,36 @@ func TestOpenAIResponsesRequestSelectsSupportedPricingModes(t *testing.T) { } } } + +func TestOpenAIResponsesContinuationSkipsPromptAffinityAndKeepsOpaqueIDs(t *testing.T) { + for _, test := range []struct { + field string + id string + err bool + }{ + {field: `"previous_response_id":" custom/id "`, id: " custom/id "}, + {field: `"previous_response_id":null`}, + {field: `"previous_response_id":""`}, + {field: `"PREVIOUS_RESPONSE_ID":"unrelated"`}, + {field: `"previous_response_id":123`, err: true}, + {field: `"previous_response_id":false`, err: true}, + {field: `"previous_response_id":[]`, err: true}, + {field: `"previous_response_id":{}`, err: true}, + {field: `"previous_response_id":"first","previous_response_id":"second"`, err: true}, + } { + t.Run(test.field, func(t *testing.T) { + request := &ParsedRequest{ + Method: http.MethodPost, Path: "/v1/responses", + Body: []byte(`{"model":"gpt-4o","input":"continue",` + test.field + `}`), + } + metadata, err := NewOpenAIResponses().InspectRequest(request) + if (err != nil) != test.err { + t.Fatalf("error = %v, want error %t", err, test.err) + } + if err == nil && (metadata.PreviousResponseID != test.id || + (len(metadata.AffinityPrefix) == 0) != (test.id != "")) { + t.Fatalf("metadata = %#v, want ID %q and mutually exclusive affinity", metadata, test.id) + } + }) + } +} diff --git a/internal/dialect/request_fields.go b/internal/dialect/request_fields.go index 4da4e2621..767ae837f 100644 --- a/internal/dialect/request_fields.go +++ b/internal/dialect/request_fields.go @@ -9,7 +9,7 @@ import ( "unicode/utf8" ) -func inspectJSONRequestFields(body []byte, requireModel bool) (RequestMetadata, error) { +func inspectJSONRequestFields(body []byte, requireModel, responsesCreate bool) (RequestMetadata, error) { if !utf8.Valid(body) { return RequestMetadata{}, fmt.Errorf("request body must be valid UTF-8") } @@ -26,6 +26,7 @@ func inspectJSONRequestFields(body []byte, requireModel bool) (RequestMetadata, result := RequestMetadata{} modelSeen := false streamSeen := false + previousResponseSeen := false for decoder.More() { fieldToken, err := decoder.Token() if err != nil { @@ -70,6 +71,14 @@ func inspectJSONRequestFields(body []byte, requireModel bool) (RequestMetadata, if !valid { return RequestMetadata{}, fmt.Errorf("stream must be a boolean") } + case responsesCreate && field == "previous_response_id": + if previousResponseSeen { + return RequestMetadata{}, fmt.Errorf("previous_response_id must be unique") + } + previousResponseSeen = true + if err := decoder.Decode(&result.PreviousResponseID); err != nil { + return RequestMetadata{}, fmt.Errorf("previous_response_id must be a string or null") + } default: var ignored json.RawMessage if err := decoder.Decode(&ignored); err != nil { diff --git a/internal/dialect/request_fields_test.go b/internal/dialect/request_fields_test.go index 42c4e639b..184aa10e7 100644 --- a/internal/dialect/request_fields_test.go +++ b/internal/dialect/request_fields_test.go @@ -40,7 +40,7 @@ func TestExtractJSONRequestFieldsStrictDecisionContract(t *testing.T) { t.Run(test.name, func(t *testing.T) { body := []byte(test.body) before := bytes.Clone(body) - metadata, err := inspectJSONRequestFields(body, true) + metadata, err := inspectJSONRequestFields(body, true, false) model := "" if metadata.Model != nil { model = *metadata.Model diff --git a/internal/dialect/rerank.go b/internal/dialect/rerank.go index 1195f3358..c08c3c7bc 100644 --- a/internal/dialect/rerank.go +++ b/internal/dialect/rerank.go @@ -31,7 +31,7 @@ func (d *Rerank) InspectRequest(request *ParsedRequest) (RequestMetadata, error) return RequestMetadata{}, fmt.Errorf("rerank requires application/json") } } - metadata, err := inspectJSONRequestFields(request.Body, true) + metadata, err := inspectJSONRequestFields(request.Body, true, false) if err != nil { return RequestMetadata{}, err } diff --git a/internal/gateway/execution_forward.go b/internal/gateway/execution_forward.go index 4315ffde8..fb61f234a 100644 --- a/internal/gateway/execution_forward.go +++ b/internal/gateway/execution_forward.go @@ -141,6 +141,17 @@ func (forwarder *ExecutionForwarder) ForwardStream( if err != nil { return false, err } + if !wasTerminal && !providerError && input.OnResponse != nil { + object, err := decodeResponsesStoreObject(event.Payload) + if err != nil { + return false, err + } + if response, exists := object["response"]; exists { + if err := input.OnResponse(response); err != nil { + return false, err + } + } + } if !wasTerminal { streamEvents.observeUsageEvent(event) if providerError { @@ -223,7 +234,7 @@ func (forwarder *ExecutionForwarder) ForwardStream( return downstreamErr } forwardData := observedData - if input.ClientProtocol == protocol.OpenAIImages || responsesStoreBuffer != nil { + if input.ClientProtocol == protocol.OpenAIImages || responsesStoreBuffer != nil || input.OnResponse != nil { forwardData = completeData if len(forwardData) == 0 { return nil diff --git a/internal/gateway/forward.go b/internal/gateway/forward.go index 3e43a11cb..78ee94fe6 100644 --- a/internal/gateway/forward.go +++ b/internal/gateway/forward.go @@ -36,6 +36,8 @@ type ForwardInput struct { UpstreamModelID string OnStreamReady func() OnFirstResponse func() + // OnResponse 在原生 Response 对象下发前登记归属,不承担上游执行。 + OnResponse func([]byte) error RequestID string AttemptID string diff --git a/internal/gateway/handler.go b/internal/gateway/handler.go index 1f9aaa297..44262888d 100644 --- a/internal/gateway/handler.go +++ b/internal/gateway/handler.go @@ -109,6 +109,7 @@ type Handler struct { routeNotFoundEvents *utils.RateLimitedEventCounter lifecycle *httplifecycle.Coordinator affinityCache *affinity.Cache + responseBindings *state.ResponseBindings } func (handler *Handler) freezeAttemptPricing( @@ -163,13 +164,14 @@ func NewHandler( manager: manager, channels: channels, subscriptions: subscriptions, registry: registry, encryption: encryptionService, forwarder: forwarder, dialects: dialects, stats: stats, mutations: mutations, limiter: limiter, requestLogSink: requestLogSink, priceTables: priceTables, - affinityCache: affinity.NewCache(), - newRequestID: newRequestID, - requestNow: time.Now, - now: time.Now, - writeTimeout: downstreamWriteTimeout, - modelListLimit: maxNonStreamingResponseBodyBytes, - logger: logrus.StandardLogger(), + affinityCache: affinity.NewCache(), + responseBindings: state.NewResponseBindings(), + newRequestID: newRequestID, + requestNow: time.Now, + now: time.Now, + writeTimeout: downstreamWriteTimeout, + modelListLimit: maxNonStreamingResponseBodyBytes, + logger: logrus.StandardLogger(), authFailureEvents: utils.NewRateLimitedEventCounter( time.Minute, time.Now, @@ -206,6 +208,7 @@ func NewHandlerWithLifecycle( priceTables PriceTableProvider, accessQuota *accessquota.Runtime, lifecycle *httplifecycle.Coordinator, + responseBindings *state.ResponseBindings, ) *Handler { handler := NewHandler( manager, @@ -227,6 +230,7 @@ func NewHandlerWithLifecycle( handler.subscriptions = subscriptions } handler.lifecycle = lifecycle + handler.responseBindings = responseBindings return handler } @@ -603,15 +607,27 @@ func (handler *Handler) Handle(ginContext *gin.Context) { allowedCredentialIDs[credentialID] = struct{}{} } query.AllowedCredentialIDs = allowedCredentialIDs - requestAffinity := handler.resolveRequestAffinity( - snapshot, - accessKey.ID, - selectedRoute.Protocol, - metadata.AffinityPrefix, - allowedCredentialRefs, - ) - query.PreferredCredentialID = requestAffinity.preferredCredentialID query.AllowedCredentialRefs = allowedCredentialRefs + var requestAffinity requestAffinity + if metadata.PreviousResponseID != "" { + binding, found := handler.responseBindings.Lookup(accessKey.ID, metadata.PreviousResponseID) + if !found { + handler.completeReason(ginContext, recorder, reasonResponseBindingNotFound) + return + } + query.AllowedCredentialIDs = map[uint]struct{}{binding.CredentialID: {}} + query.AllowedCredentialRefs = map[uint]state.CredentialRef{ + binding.CredentialID: { + ID: binding.CredentialID, GroupID: binding.GroupID, + IdentityGeneration: binding.IdentityGeneration, + }, + } + } else { + requestAffinity = handler.resolveRequestAffinity( + snapshot, accessKey.ID, selectedRoute.Protocol, metadata.AffinityPrefix, allowedCredentialRefs, + ) + query.PreferredCredentialID = requestAffinity.preferredCredentialID + } iterator := scheduler.New(snapshot, handler.registry, query) handler.executeAttempts( ginContext, @@ -908,8 +924,8 @@ func (handler *Handler) executeAttempts( scope execution.ErrorScope, ) bool { attemptSequence++ - if attemptSequence == 1 && requestAffinity.preferredCredentialID != 0 && - selection.CredentialID == requestAffinity.preferredCredentialID { + if attemptSequence == 1 && (originalMetadata.PreviousResponseID != "" || + (requestAffinity.preferredCredentialID != 0 && selection.CredentialID == requestAffinity.preferredCredentialID)) { recorder.setAffinityHit(true) } updateDebugHeaders(ginContext.Writer.Header(), selection.Group.Name, attemptSequence) @@ -1119,8 +1135,8 @@ func (handler *Handler) executeAttempts( attemptSequence++ forwardAttempts++ - if attemptSequence == 1 && requestAffinity.preferredCredentialID != 0 && - selection.CredentialID == requestAffinity.preferredCredentialID { + if attemptSequence == 1 && (originalMetadata.PreviousResponseID != "" || + (requestAffinity.preferredCredentialID != 0 && selection.CredentialID == requestAffinity.preferredCredentialID)) { recorder.setAffinityHit(true) } updateDebugHeaders(ginContext.Writer.Header(), selection.Group.Name, attemptSequence) @@ -1154,6 +1170,7 @@ func (handler *Handler) executeAttempts( ProxyFingerprint: proxyFingerprint, ForceCredentialRefresh: forceCredentialRefresh, ContinuityKey: requestAffinity.continuityKey, + OnResponse: handler.responseBindingObserver(recorder.accessKeyID, selection, ref, prepared.request), OnFirstResponse: func() { recorder.recordFirstResponse() }, @@ -1176,6 +1193,20 @@ func (handler *Handler) executeAttempts( result = handler.forwarder.Forward(ginContext.Request.Context(), input) } result = normalizeUpstreamResultContract(result) + if !stream && result.HasResponse() && !result.ProviderErrorBeforeCommit && + result.DispatchState != execution.DispatchLocal && + result.StatusCode >= http.StatusOK && result.StatusCode < http.StatusMultipleChoices { + if input.OnResponse != nil { + if err := input.OnResponse(result.Body); err != nil { + result.Err = err + result.ExecutionError = &execution.ErrorEvidence{ + Kind: execution.ErrorKindInternal, OriginHint: execution.ErrorOriginInternal, + ScopeHint: execution.ErrorScopeRequest, Code: "response_binding_conflict", + Summary: "Response ownership could not be recorded.", ReplaySafety: execution.ReplaySafetyUnknown, + } + } + } + } attemptCompleted := time.Time{} if recorder != nil { attemptCompleted = recorder.now() @@ -1216,7 +1247,9 @@ func (handler *Handler) executeAttempts( ) if stream && result.Stream.EndReason == StreamEndCleanEOF { handler.recordCredentialSuccess(ref, attemptNow) - handler.recordAffinitySuccess(requestAffinity, selection, ref) + if originalMetadata.PreviousResponseID == "" { + handler.recordAffinitySuccess(requestAffinity, selection, ref) + } } return } @@ -1313,7 +1346,8 @@ func (handler *Handler) executeAttempts( return } if result.DispatchState != execution.DispatchLocal && - result.StatusCode >= http.StatusOK && result.StatusCode < http.StatusMultipleChoices { + result.StatusCode >= http.StatusOK && result.StatusCode < http.StatusMultipleChoices && + originalMetadata.PreviousResponseID == "" { handler.recordAffinitySuccess(requestAffinity, selection, ref) } return diff --git a/internal/gateway/reason.go b/internal/gateway/reason.go index a6cc586e0..1bf471364 100644 --- a/internal/gateway/reason.go +++ b/internal/gateway/reason.go @@ -33,6 +33,7 @@ var ( reasonEndpointNotFound = reason{Status: http.StatusNotFound, Code: "protocol_endpoint_not_found", Message: "Protocol endpoint not found."} reasonMethodNotAllowed = reason{Status: http.StatusMethodNotAllowed, Code: "method_not_allowed", Message: "Method not allowed."} reasonInvalidProtocolRequest = reason{Status: http.StatusBadRequest, Code: "invalid_protocol_request", Message: "Invalid protocol request."} + reasonResponseBindingNotFound = reason{Status: http.StatusBadRequest, Code: "response_binding_not_found", Message: "Previous response ownership could not be located."} reasonModelRequiredByFilter = reason{Status: http.StatusBadRequest, Code: "model_required_by_filter", Message: "A model is required by the access key filter."} reasonNoCandidate = reason{Status: http.StatusServiceUnavailable, Code: "no_available_candidate", Message: "No available upstream candidate."} reasonUpstreamRateLimited = reason{Status: http.StatusTooManyRequests, Code: "upstream_rate_limited", Message: "Upstream rate limit exceeded."} diff --git a/internal/gateway/reason_test.go b/internal/gateway/reason_test.go index 8f401ebc5..6371ee946 100644 --- a/internal/gateway/reason_test.go +++ b/internal/gateway/reason_test.go @@ -14,6 +14,7 @@ func TestWriteReasonUsesStableDataPlaneEnvelope(t *testing.T) { reasonInvalidAccessKey, reasonEndpointNotFound, reasonInvalidProtocolRequest, + reasonResponseBindingNotFound, reasonModelRequiredByFilter, reasonNoCandidate, reasonUpstreamConnect, diff --git a/internal/gateway/responses_continuation.go b/internal/gateway/responses_continuation.go new file mode 100644 index 000000000..6c14e6e0c --- /dev/null +++ b/internal/gateway/responses_continuation.go @@ -0,0 +1,46 @@ +package gateway + +import ( + "encoding/json" + "fmt" + "net/http" + + "gpt-load/internal/dialect" + "gpt-load/internal/execution" + "gpt-load/internal/scheduler" + "gpt-load/internal/state" +) + +func (handler *Handler) responseBindingObserver( + accessKeyID uint, + selection scheduler.Selection, + ref state.CredentialRef, + request *dialect.ParsedRequest, +) func([]byte) error { + if request == nil || request.Method != http.MethodPost || request.Path != "/v1/responses" || + selection.RouteMode != execution.RouteNative || selection.ResponsesStoreDowngraded || + !selection.ResolvedTarget.SupportsResponsesLifecycle() { + return nil + } + var options struct { + Store *bool `json:"store"` + } + if json.Unmarshal(request.Body, &options) != nil || (options.Store != nil && !*options.Store) { + return nil + } + return func(payload []byte) error { + var response struct { + ID string `json:"id"` + Object string `json:"object"` + Store *bool `json:"store"` + } + if json.Unmarshal(payload, &response) != nil || response.Object != "response" || + response.ID == "" || (response.Store != nil && !*response.Store) { + return nil + } + if !handler.responseBindings.Record(accessKeyID, response.ID, ref) { + return fmt.Errorf("%w: response ownership could not be recorded", ErrUpstreamProtocol) + } + return nil + } +} diff --git a/internal/gateway/responses_continuation_test.go b/internal/gateway/responses_continuation_test.go new file mode 100644 index 000000000..a3e222495 --- /dev/null +++ b/internal/gateway/responses_continuation_test.go @@ -0,0 +1,307 @@ +package gateway + +import ( + "bytes" + "compress/gzip" + "context" + "fmt" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/gin-gonic/gin" + + "gpt-load/internal/app" + "gpt-load/internal/dialect" + "gpt-load/internal/execution" + "gpt-load/internal/parameteroverride" + "gpt-load/internal/state" +) + +func TestResponsesContinuationPinsCredentialWithoutSoftAffinity(t *testing.T) { + for _, enabled := range []bool{false, true} { + t.Run(fmt.Sprint(enabled), func(t *testing.T) { + forwarder := &scriptedForwarder{results: []UpstreamResult{ + storedResponse("first"), storedResponse("second"), storedResponse("third"), + }} + handler, engine, sink := newContinuationFixture(t, forwarder) + group := handler.manager.Current().Groups[1] + group.AffinityEnabled = enabled + handler.manager.Current().Groups[1] = group + handler.manager.Current().Settings.AffinityEnabled = enabled + + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","input":"initial","store":true}`, http.StatusOK) + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","input":"continue","previous_response_id":"first"}`, http.StatusOK) + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","input":"continue"}`, http.StatusOK) + + assertAffinityAttemptKeys(t, forwarder.inputs, []string{"sk-one", "sk-one", "sk-two"}) + assertAffinityHits(t, sink.snapshot(), []bool{false, true, false}) + if got := handler.registry.SchedulingState().CaptureCheckpoint().Sequence; got != 3 { + t.Fatalf("scheduling allocations = %d, want 3", got) + } + }) + } +} + +func TestResponsesContinuationRejectsUnknownAndOtherAccessKeyIDs(t *testing.T) { + forwarder := &scriptedForwarder{results: []UpstreamResult{storedResponse("first")}} + handler, engine, sink := newContinuationFixture(t, forwarder) + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","input":"initial"}`, http.StatusOK) + other := handler.manager.Current().AccessKeysByHash[handler.encryption.Hash("gl-client")] + other.ID = 2 + handler.manager.Current().AccessKeysByHash[handler.encryption.Hash("gl-other")] = other + handler.manager.Current().AccessKeysByID[other.ID] = other + for _, request := range []struct{ key, id string }{{"gl-other", "first"}, {"gl-client", "unknown"}} { + serveContinuation(t, engine, request.key, fmt.Sprintf(`{"model":"gpt-4o","input":"continue","previous_response_id":%q}`, request.id), http.StatusBadRequest) + } + if len(forwarder.inputs) != 1 { + t.Fatalf("upstream attempts = %d, want only the initial request", len(forwarder.inputs)) + } + assertAffinityHits(t, sink.snapshot(), []bool{false, false, false}) +} + +func TestResponsesContinuationUsesCurrentCandidateIntersection(t *testing.T) { + forwarder := &scriptedForwarder{results: []UpstreamResult{storedResponse("first")}} + handler, engine, sink := newContinuationFixture(t, forwarder) + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","input":"initial"}`, http.StatusOK) + if err := handler.registry.(*state.CredentialRegistry).SetCredentialStatus(1, state.CredentialStatusDisabled); err != nil { + t.Fatal(err) + } + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","input":"continue","previous_response_id":"first"}`, http.StatusServiceUnavailable) + if len(forwarder.inputs) != 1 { + t.Fatal("continuation escaped to the other credential") + } + assertAffinityHits(t, sink.snapshot(), []bool{false, false}) +} + +func TestResponsesContinuationRegistersSSEBeforeDelivery(t *testing.T) { + var engine *gin.Engine + writer := httptest.NewRecorder() + var credentials []uint + created := "event: response.created\r\n" + `data: {"type":"response.created","response":{"id":"first","object":"response","store":true}}` + "\r\n\r\n" + completed := "event: response.completed\r\n" + `data: {"type":"response.completed","response":{"id":"first","object":"response","store":true}}` + "\r\n\r\n" + executor := fakeExecutionExecutor{ + unary: func(_ context.Context, spec execution.AttemptSpec) execution.AttemptResult { + credentials = append(credentials, spec.Credential.ID) + return execution.AttemptResult{ + StatusCode: http.StatusOK, DispatchState: execution.DispatchMaybeSent, ResponseStarted: true, + Header: http.Header{"Content-Type": {"application/json"}}, Body: storedResponse("second").Body, + } + }, + stream: func(_ context.Context, spec execution.AttemptSpec, sink execution.StreamSink) execution.StreamResult { + credentials = append(credentials, spec.Credential.ID) + for _, event := range []execution.StreamEvent{ + {Sequence: 1, Kind: execution.StreamEventReady, StatusCode: http.StatusOK, Header: http.Header{"Content-Type": {"text/event-stream"}}}, + {Sequence: 2, Kind: execution.StreamEventData, Data: []byte(created[:len(created)-4])}, + } { + if err := sink(event); err != nil { + t.Fatal(err) + } + } + if writer.Body.Len() != 0 { + t.Fatal("partial response ID was delivered before registration") + } + if err := sink(execution.StreamEvent{Sequence: 3, Kind: execution.StreamEventData, Data: []byte(created[len(created)-4:])}); err != nil { + t.Fatal(err) + } + if writer.Body.String() != created { + t.Fatalf("created event was not delivered intact: %q", writer.Body.String()) + } + // 客户端收到 created 即可续接,无需等待原请求 completed。 + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","previous_response_id":"first","input":"continue"}`, http.StatusOK) + if err := sink(execution.StreamEvent{Sequence: 4, Kind: execution.StreamEventData, Data: []byte(completed)}); err != nil { + t.Fatal(err) + } + return execution.StreamResult{StatusCode: http.StatusOK, DispatchState: execution.DispatchMaybeSent, ResponseStarted: true} + }, + } + _, engine, _ = newContinuationFixture(t, NewExecutionForwarder(executor)) + request := httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewBufferString(`{"model":"gpt-4o","input":"initial","stream":true}`)) + request.Header.Set("Authorization", "Bearer gl-client") + engine.ServeHTTP(writer, request) + if writer.Body.String() != created+completed || fmt.Sprint(credentials) != "[1 1]" { + t.Fatalf("SSE response = %q, credentials = %v", writer.Body.String(), credentials) + } +} + +func TestResponsesContinuationKeepsBindingAcrossModelCooldownAndConfigPublication(t *testing.T) { + forwarder := &scriptedForwarder{results: []UpstreamResult{storedResponse("first"), storedResponse("second")}} + handler, engine, sink := newContinuationFixture(t, forwarder) + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","input":"initial"}`, http.StatusOK) + ref, _ := handler.registry.CredentialRef(1) + registry := handler.registry.(*state.CredentialRegistry) + if ok, _ := registry.SetModelCooldown(ref, "gpt-4o", time.Now().Add(time.Hour), time.Now()); !ok { + t.Fatal("set model cooldown failed") + } + response := serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","previous_response_id":"first","input":"continue"}`, http.StatusTooManyRequests) + if response.Header().Get("Retry-After") == "" || len(forwarder.inputs) != 1 { + t.Fatal("cooldown did not reuse existing candidate rejection") + } + registry.ClearModelCooldowns(1) + handler.manager.Current().Revision++ + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","previous_response_id":"first","input":"continue"}`, http.StatusOK) + assertAffinityAttemptKeys(t, forwarder.inputs, []string{"sk-one", "sk-one"}) + assertAffinityHits(t, sink.snapshot(), []bool{false, false, true}) +} + +func TestResponsesContinuationFailureExhaustsBoundCandidate(t *testing.T) { + forwarder := &scriptedForwarder{results: []UpstreamResult{ + storedResponse("first"), + {StatusCode: http.StatusServiceUnavailable, Body: []byte(`{"error":{"message":"upstream unavailable"}}`)}, + storedResponse("must-not-retry"), + }} + handler, engine, sink := newContinuationFixture(t, forwarder) + handler.manager.Current().Settings.RetryCount = 3 + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","input":"initial"}`, http.StatusOK) + response := serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","previous_response_id":"first","input":"continue"}`, http.StatusServiceUnavailable) + if response.Body.String() != string(forwarder.results[1].Body) { + t.Fatal("the original upstream error was replaced") + } + assertAffinityAttemptKeys(t, forwarder.inputs, []string{"sk-one", "sk-one"}) + events := sink.snapshot() + if len(events[1].Attempts) != 1 || events[1].Attempts[0].WillRetry || !events[1].AffinityHit { + t.Fatalf("failure telemetry = %#v", events[1]) + } +} + +func TestResponsesContinuationRejectsChangedCredentialIdentity(t *testing.T) { + forwarder := &scriptedForwarder{results: []UpstreamResult{storedResponse("first")}} + handler, engine, _ := newContinuationFixture(t, forwarder) + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","input":"initial"}`, http.StatusOK) + ref, _ := handler.registry.CredentialRef(1) + registry := handler.registry.(*state.CredentialRegistry) + if err := registry.ApplyCredentialImport(1, []state.CredentialEntry{{ + ID: 1, GroupID: 1, Status: state.CredentialStatusActive, Version: ref.Version + 1, + IdentityGeneration: ref.IdentityGeneration + 10, Fingerprint: ref.Fingerprint, EncryptedValue: ref.EncryptedValue, + }}); err != nil { + t.Fatal(err) + } + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","previous_response_id":"first","input":"continue"}`, http.StatusServiceUnavailable) + if len(forwarder.inputs) != 1 { + t.Fatal("old ID inherited a replacement identity") + } +} + +func TestResponsesContinuationDoesNotLearnUnstoredOrUnrelatedResponses(t *testing.T) { + for _, test := range []struct{ request, response string }{ + {`{"model":"gpt-4o","input":"initial","store":false}`, string(storedResponse("first").Body)}, + {`{"model":"gpt-4o","input":"initial"}`, `{"id":"first","object":"response","store":false}`}, + {`{"model":"gpt-4o","input":"initial"}`, `{"id":"first","object":"chat.completion"}`}, + } { + forwarder := &scriptedForwarder{results: []UpstreamResult{{StatusCode: http.StatusOK, Body: []byte(test.response)}}} + _, engine, _ := newContinuationFixture(t, forwarder) + serveContinuation(t, engine, "gl-client", test.request, http.StatusOK) + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","previous_response_id":"first","input":"continue"}`, http.StatusBadRequest) + if len(forwarder.inputs) != 1 { + t.Fatal("unavailable response ID reached an upstream") + } + } +} + +func TestResponsesContinuationRejectsConflictingResponseBeforeDelivery(t *testing.T) { + forwarder := &scriptedForwarder{results: []UpstreamResult{ + storedResponse("shared-id"), storedResponse("shared-id"), storedResponse("continued"), + }} + _, engine, _ := newContinuationFixture(t, forwarder) + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","input":"first"}`, http.StatusOK) + response := serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","input":"independent"}`, http.StatusBadGateway) + if bytes.Contains(response.Body.Bytes(), []byte("shared-id")) { + t.Fatal("conflicting response ID was exposed") + } + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","previous_response_id":"shared-id","input":"continue"}`, http.StatusOK) + assertAffinityAttemptKeys(t, forwarder.inputs, []string{"sk-one", "sk-two", "sk-one"}) +} + +func TestResponsesContinuationRejectsLegacyOverrideBeforeForward(t *testing.T) { + forwarder := &scriptedForwarder{results: []UpstreamResult{storedResponse("first")}} + handler, engine, _ := newContinuationFixture(t, forwarder) + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","input":"initial"}`, http.StatusOK) + rules, err := parameteroverride.Compile([]any{map[string]any{"set": map[string]any{"previous_response_id": "wrong"}}}) + if err != nil { + t.Fatal(err) + } + group := handler.manager.Current().Groups[1] + group.ParameterOverrides = rules + handler.manager.Current().Groups[1] = group + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","previous_response_id":"first","input":"continue"}`, http.StatusServiceUnavailable) + if len(forwarder.inputs) != 1 { + t.Fatal("legacy override dispatched a changed continuation") + } +} + +func TestResponsesContinuationResumesFromRuntimeCheckpoint(t *testing.T) { + before, engine, _ := newContinuationFixture(t, &scriptedForwarder{results: []UpstreamResult{storedResponse("before-restart")}}) + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","input":"initial"}`, http.StatusOK) + dir := t.TempDir() + checkpoint := app.NewFileRuntimeStateCheckpoint(dir, nil, nil, before.responseBindings) + if err := checkpoint.Save(context.Background()); err != nil { + t.Fatal(err) + } + forwarder := &scriptedForwarder{results: []UpstreamResult{storedResponse("new-root"), storedResponse("continued")}} + after, restarted, _ := newContinuationFixture(t, forwarder) + if err := app.NewFileRuntimeStateCheckpoint(dir, nil, nil, after.responseBindings).Restore(context.Background()); err != nil { + t.Fatal(err) + } + serveContinuation(t, restarted, "gl-client", `{"model":"gpt-4o","input":"another root"}`, http.StatusOK) + serveContinuation(t, restarted, "gl-client", `{"model":"gpt-4o","previous_response_id":"before-restart","input":"continue"}`, http.StatusOK) + assertAffinityAttemptKeys(t, forwarder.inputs, []string{"sk-one", "sk-one"}) +} + +func TestResponsesContinuationLearnsCompressedJSONResponse(t *testing.T) { + var compressed bytes.Buffer + zipper := gzip.NewWriter(&compressed) + if _, err := zipper.Write(storedResponse("compressed-response").Body); err != nil { + t.Fatal(err) + } + if err := zipper.Close(); err != nil { + t.Fatal(err) + } + var credentials []uint + executor := fakeExecutionExecutor{unary: func(_ context.Context, spec execution.AttemptSpec) execution.AttemptResult { + credentials = append(credentials, spec.Credential.ID) + return execution.AttemptResult{ + StatusCode: http.StatusOK, DispatchState: execution.DispatchMaybeSent, ResponseStarted: true, + Header: http.Header{"Content-Type": {"application/json"}, "Content-Encoding": {"gzip"}}, + Body: compressed.Bytes(), + } + }} + _, engine, _ := newContinuationFixture(t, NewExecutionForwarder(executor)) + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","input":"initial"}`, http.StatusOK) + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","previous_response_id":"compressed-response","input":"continue"}`, http.StatusOK) + if fmt.Sprint(credentials) != "[1 1]" { + t.Fatalf("compressed response resumed through credentials %v", credentials) + } +} + +func storedResponse(id string) UpstreamResult { + return UpstreamResult{ + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": {"application/json"}}, + Body: []byte(fmt.Sprintf(`{"id":%q,"object":"response","status":"completed","store":true,"output":[]}`, id)), + } +} + +func newContinuationFixture(t *testing.T, forwarder AttemptForwarder) (*Handler, *gin.Engine, *recordingRequestLogSink) { + t.Helper() + handler, _, _ := newHandlerForTest(t, forwarder, "sk-one", "sk-two") + handler.dialects = dialect.NewSet(dialect.NewOpenAIResponses()) + sink := &recordingRequestLogSink{} + handler.requestLogSink = sink + engine := gin.New() + bindGatewayRoutesForTest(t, engine, handler) + return handler, engine, sink +} + +func serveContinuation(t *testing.T, engine http.Handler, key, body string, status int) *httptest.ResponseRecorder { + t.Helper() + request := httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewBufferString(body)) + request.Header.Set("Authorization", "Bearer "+key) + request.Header.Set("Content-Type", "application/json") + response := httptest.NewRecorder() + engine.ServeHTTP(response, request) + if response.Code != status { + t.Fatalf("status = %d, want %d; body = %s", response.Code, status, response.Body.String()) + } + return response +} diff --git a/internal/parameteroverride/response_continuation_test.go b/internal/parameteroverride/response_continuation_test.go new file mode 100644 index 000000000..4678d7e73 --- /dev/null +++ b/internal/parameteroverride/response_continuation_test.go @@ -0,0 +1,33 @@ +package parameteroverride + +import ( + "testing" + + "gpt-load/internal/execution" + "gpt-load/internal/protocol" +) + +func TestResponsesContinuationOverrideLoadsButFailsOnlyMatchingRequests(t *testing.T) { + for _, action := range []map[string]any{ + {"set": map[string]any{"previous_response_id": "replacement"}}, + {"remove": []any{"/previous_response_id"}}, + {"remove": []any{"/previous_response_id/value"}}, + } { + action["match"] = map[string]any{"model": "matched"} + rules, err := Compile([]any{action}) + if err != nil { + t.Fatalf("legacy rules must still compile: %v", err) + } + body := []byte(`{"model":"matched","previous_response_id":"original"}`) + if _, _, err := rules.Apply(protocol.OpenAIResponses, execution.OperationResponsesCreate, "matched", body); err == nil { + t.Fatal("matching override changed the continuation ID") + } + got, applied, err := rules.Apply(protocol.OpenAIResponses, execution.OperationResponsesCreate, "unmatched", body) + if err != nil || applied || string(got) != string(body) { + t.Fatalf("unmatched rule affected request: %s %t %v", got, applied, err) + } + if _, _, err := rules.Apply(protocol.OpenAICompletions, execution.OperationChatCompletion, "matched", body); err != nil { + t.Fatalf("Responses restriction changed another protocol: %v", err) + } + } +} diff --git a/internal/parameteroverride/rules.go b/internal/parameteroverride/rules.go index 1e9c23c79..faa00b434 100644 --- a/internal/parameteroverride/rules.go +++ b/internal/parameteroverride/rules.go @@ -197,6 +197,32 @@ func compileMatch(raw json.RawMessage) (compiledMatch, error) { // Empty reports whether no rule can be applied. func (rules Rules) Empty() bool { return len(rules.entries) == 0 } +// ValidateResponsesContinuation 用于管理面保存;不改变历史配置的 Compile 行为。 +func (rules Rules) ValidateResponsesContinuation() error { + for _, entry := range rules.entries { + if entry.clientProtocol == "" || entry.clientProtocol == protocol.OpenAIResponses { + if err := entry.validateResponsesContinuation(); err != nil { + return err + } + } + } + return nil +} + +func (entry rule) validateResponsesContinuation() error { + for key := range entry.set { + if strings.EqualFold(key, "previous_response_id") { + return fmt.Errorf("parameter overrides cannot change previous_response_id") + } + } + for _, path := range entry.remove { + if strings.EqualFold(path[0], "previous_response_id") { + return fmt.Errorf("parameter overrides cannot change previous_response_id") + } + } + return nil +} + // Clone returns an independently owned rule set. func (rules Rules) Clone() Rules { cloned := Rules{entries: make([]rule, len(rules.entries))} @@ -225,6 +251,11 @@ func (rules Rules) Apply( matched := make([]rule, 0, len(rules.entries)) for _, entry := range rules.entries { if entry.matches(clientProtocol, clientModel) { + if clientProtocol == protocol.OpenAIResponses { + if err := entry.validateResponsesContinuation(); err != nil { + return nil, false, err + } + } matched = append(matched, entry) } } diff --git a/internal/state/response_bindings.go b/internal/state/response_bindings.go new file mode 100644 index 000000000..7a04b065e --- /dev/null +++ b/internal/state/response_bindings.go @@ -0,0 +1,158 @@ +package state + +import ( + "container/list" + "fmt" + "sort" + "sync" + "time" +) + +const ( + DefaultResponseBindingTTL = 30 * 24 * time.Hour + DefaultResponseBindingCapacity = 100_000 + maxResponseBindingIDBytes = 16 << 20 +) + +// ResponseBinding 只保存响应归属;可用性仍由当前路由和凭据运行态决定。 +type ResponseBinding struct { + AccessKeyID uint `json:"access_key_id"` + ResponseID string `json:"response_id"` + GroupID uint `json:"group_id"` + CredentialID uint `json:"credential_id"` + IdentityGeneration uint64 `json:"identity_generation"` + ExpiresAt time.Time `json:"expires_at"` +} + +type responseBindingKey struct { + accessKeyID uint + responseID string +} + +// ResponseBindings 是有界的内存归属索引,不持有 DB、文件或软亲和配置。 +type ResponseBindings struct { + mu sync.Mutex + entries map[responseBindingKey]*list.Element + order list.List + idBytes int + capacity int + ttl time.Duration + now func() time.Time +} + +func NewResponseBindings() *ResponseBindings { + return &ResponseBindings{ + entries: make(map[responseBindingKey]*list.Element), + capacity: DefaultResponseBindingCapacity, + ttl: DefaultResponseBindingTTL, + now: time.Now, + } +} + +func (bindings *ResponseBindings) Lookup(accessKeyID uint, responseID string) (ResponseBinding, bool) { + if bindings == nil { + return ResponseBinding{}, false + } + bindings.mu.Lock() + defer bindings.mu.Unlock() + element := bindings.entries[responseBindingKey{accessKeyID, responseID}] + if element == nil { + return ResponseBinding{}, false + } + binding := element.Value.(ResponseBinding) + if !binding.ExpiresAt.After(bindings.now()) { + bindings.remove(element) + return ResponseBinding{}, false + } + return binding, true +} + +// Record 在响应下发前登记;不同归属冲突时拒绝当前响应,不覆盖已有归属。 +func (bindings *ResponseBindings) Record(accessKeyID uint, responseID string, ref CredentialRef) bool { + if bindings == nil || accessKeyID == 0 || responseID == "" || ref.ID == 0 || + ref.GroupID == 0 || ref.IdentityGeneration == 0 || len(responseID) > maxResponseBindingIDBytes { + return false + } + bindings.mu.Lock() + defer bindings.mu.Unlock() + now := bindings.now() + bindings.expire(now) + return bindings.insert(ResponseBinding{ + AccessKeyID: accessKeyID, ResponseID: responseID, + GroupID: ref.GroupID, CredentialID: ref.ID, IdentityGeneration: ref.IdentityGeneration, + ExpiresAt: now.Add(bindings.ttl), + }) +} + +func (bindings *ResponseBindings) insert(binding ResponseBinding) bool { + key := responseBindingKey{binding.AccessKeyID, binding.ResponseID} + if element := bindings.entries[key]; element != nil { + existing := element.Value.(ResponseBinding) + return existing.GroupID == binding.GroupID && existing.CredentialID == binding.CredentialID && + existing.IdentityGeneration == binding.IdentityGeneration + } + if bindings.capacity <= 0 || bindings.ttl <= 0 { + return false + } + for len(bindings.entries) >= bindings.capacity || bindings.idBytes+len(binding.ResponseID) > maxResponseBindingIDBytes { + bindings.remove(bindings.order.Front()) + } + bindings.entries[key] = bindings.order.PushBack(binding) + bindings.idBytes += len(binding.ResponseID) + return true +} + +func (bindings *ResponseBindings) CaptureCheckpoint() []ResponseBinding { + bindings.mu.Lock() + defer bindings.mu.Unlock() + bindings.expire(bindings.now()) + checkpoint := make([]ResponseBinding, 0, len(bindings.entries)) + for element := bindings.order.Front(); element != nil; element = element.Next() { + checkpoint = append(checkpoint, element.Value.(ResponseBinding)) + } + return checkpoint +} + +// RestoreCheckpoint 只恢复归属,不在这里复制路由、权限和健康判断。 +func (bindings *ResponseBindings) RestoreCheckpoint(checkpoint []ResponseBinding) error { + ordered := append([]ResponseBinding(nil), checkpoint...) + sort.SliceStable(ordered, func(i, j int) bool { return ordered[i].ExpiresAt.Before(ordered[j].ExpiresAt) }) + bindings.mu.Lock() + defer bindings.mu.Unlock() + bindings.reset() + now := bindings.now() + for _, binding := range ordered { + if binding.AccessKeyID == 0 || binding.ResponseID == "" || binding.CredentialID == 0 || + binding.GroupID == 0 || binding.IdentityGeneration == 0 || !binding.ExpiresAt.After(now) || + len(binding.ResponseID) > maxResponseBindingIDBytes { + continue + } + if !bindings.insert(binding) { + bindings.reset() + return fmt.Errorf("response checkpoint contains conflicting ownership") + } + } + return nil +} + +func (bindings *ResponseBindings) reset() { + bindings.entries = make(map[responseBindingKey]*list.Element) + bindings.order.Init() + bindings.idBytes = 0 +} + +func (bindings *ResponseBindings) expire(now time.Time) { + for element := bindings.order.Front(); element != nil; element = bindings.order.Front() { + if element.Value.(ResponseBinding).ExpiresAt.After(now) { + return + } + bindings.remove(element) + } +} + +func (bindings *ResponseBindings) remove(element *list.Element) { + binding := element.Value.(ResponseBinding) + delete(bindings.entries, responseBindingKey{binding.AccessKeyID, binding.ResponseID}) + bindings.idBytes -= len(binding.ResponseID) + bindings.order.Remove(element) +} diff --git a/internal/state/response_bindings_test.go b/internal/state/response_bindings_test.go new file mode 100644 index 000000000..1b6a124c0 --- /dev/null +++ b/internal/state/response_bindings_test.go @@ -0,0 +1,128 @@ +package state + +import ( + "strings" + "sync" + "testing" + "time" +) + +func TestResponseBindingsPreserveIdentityAndOriginalExpiry(t *testing.T) { + now := time.Date(2026, 9, 9, 0, 0, 0, 0, time.UTC) + bindings := NewResponseBindings() + bindings.now = func() time.Time { return now } + bindings.ttl = time.Hour + ref := CredentialRef{ID: 1, GroupID: 2, IdentityGeneration: 3, Version: 1} + if !bindings.Record(1, "opaque/id", ref) { + t.Fatal("record failed") + } + now = now.Add(30 * time.Minute) + ref.Version++ + if !bindings.Record(1, "opaque/id", ref) { + t.Fatal("normal token refresh changed ownership") + } + ref.IdentityGeneration++ + if bindings.Record(1, "opaque/id", ref) { + t.Fatal("another identity replaced existing ownership") + } + if saved, ok := bindings.Lookup(1, "opaque/id"); !ok || saved.IdentityGeneration != 3 { + t.Fatal("ownership changed after conflict") + } + if _, ok := bindings.Lookup(2, "opaque/id"); ok { + t.Fatal("binding crossed AccessKeys") + } + now = now.Add(30 * time.Minute) + if _, ok := bindings.Lookup(1, "opaque/id"); ok { + t.Fatal("duplicate observation extended the original expiry") + } +} + +func TestResponseBindingsEvictOldestAndBoundIDMemory(t *testing.T) { + bindings := NewResponseBindings() + bindings.capacity = 2 + ref := CredentialRef{ID: 1, GroupID: 1, IdentityGeneration: 1} + for _, id := range []string{"old", "new", "newest"} { + if !bindings.Record(1, id, ref) { + t.Fatal("record failed") + } + } + if _, ok := bindings.Lookup(1, "old"); ok { + t.Fatal("capacity retained oldest entry") + } + for _, id := range []string{"new", "newest"} { + if _, ok := bindings.Lookup(1, id); !ok { + t.Fatalf("recent binding %q was lost", id) + } + } + largeID := strings.Repeat("x", maxResponseBindingIDBytes) + if !bindings.Record(1, largeID, ref) { + t.Fatal("opaque ID within memory budget was rejected") + } + if _, ok := bindings.Lookup(1, "newest"); ok { + t.Fatal("ID byte budget did not evict old entries") + } + if bindings.Record(2, largeID+"x", ref) { + t.Fatal("oversized record exceeded ID memory budget") + } +} + +func TestResponseBindingsConcurrentRegistrationIsConsistent(t *testing.T) { + bindings := NewResponseBindings() + var workers sync.WaitGroup + for range 16 { + workers.Go(func() { + for range 100 { + if !bindings.Record(1, "same-response", CredentialRef{ID: 1, GroupID: 1, IdentityGeneration: 1}) { + t.Error("idempotent concurrent record failed") + } + if _, ok := bindings.Lookup(1, "same-response"); !ok { + t.Error("registered response was lost") + } + } + }) + } + workers.Wait() +} + +func TestResponseBindingsCheckpointKeepsExpiryAndCapacity(t *testing.T) { + now := time.Date(2026, 9, 9, 0, 0, 0, 0, time.UTC) + original := NewResponseBindings() + original.now = func() time.Time { return now } + original.ttl = time.Hour + ref := CredentialRef{ID: 1, GroupID: 1, IdentityGeneration: 1} + original.Record(1, "expired", ref) + now = now.Add(30 * time.Minute) + original.Record(1, "retained", ref) + checkpoint := original.CaptureCheckpoint() + now = now.Add(30 * time.Minute) + restored := NewResponseBindings() + restored.capacity = 1 + restored.now = func() time.Time { return now } + if err := restored.RestoreCheckpoint(checkpoint); err != nil { + t.Fatal(err) + } + if _, ok := restored.Lookup(1, "expired"); ok { + t.Fatal("expired checkpoint binding was restored") + } + if binding, ok := restored.Lookup(1, "retained"); !ok || !binding.ExpiresAt.Equal(checkpoint[1].ExpiresAt) { + t.Fatal("restore lost or extended retained binding") + } + now = now.Add(30 * time.Minute) + if _, ok := restored.Lookup(1, "retained"); ok { + t.Fatal("restart renewed response TTL") + } +} + +func TestResponseBindingsRejectConflictingCheckpoint(t *testing.T) { + bindings := NewResponseBindings() + first := ResponseBinding{AccessKeyID: 1, ResponseID: "same", CredentialID: 1, GroupID: 1, + IdentityGeneration: 1, ExpiresAt: time.Now().Add(time.Hour)} + conflict := first + conflict.CredentialID = 2 + if err := bindings.RestoreCheckpoint([]ResponseBinding{first, conflict}); err == nil { + t.Fatal("ambiguous checkpoint was accepted") + } + if _, ok := bindings.Lookup(1, "same"); ok { + t.Fatal("conflicting checkpoint left routable ownership") + } +} From 7f28bbd3227101ac6ae9fea81dd48d7de2faf144 Mon Sep 17 00:00:00 2001 From: tbphp Date: Wed, 9 Sep 2026 21:06:42 +0800 Subject: [PATCH 2/3] =?UTF-8?q?fix(responses):=20=E9=99=90=E5=88=B6?= =?UTF-8?q?=E5=93=8D=E5=BA=94=20ID=20=E9=95=BF=E5=BA=A6=E5=B9=B6=E5=85=BC?= =?UTF-8?q?=E5=AE=B9=E7=BC=BA=E5=A4=B1=E5=AD=97=E6=AE=B5=E5=88=A0=E9=99=A4?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../response_continuation_test.go | 39 ++++++++++++ internal/parameteroverride/rules.go | 25 +++++--- internal/state/response_bindings.go | 5 +- internal/state/response_bindings_test.go | 59 ++++++++++++++++--- 4 files changed, 112 insertions(+), 16 deletions(-) diff --git a/internal/parameteroverride/response_continuation_test.go b/internal/parameteroverride/response_continuation_test.go index 4678d7e73..6e4f08474 100644 --- a/internal/parameteroverride/response_continuation_test.go +++ b/internal/parameteroverride/response_continuation_test.go @@ -31,3 +31,42 @@ func TestResponsesContinuationOverrideLoadsButFailsOnlyMatchingRequests(t *testi } } } + +func TestResponsesContinuationOverrideAllowsOnlyMissingRemovals(t *testing.T) { + for _, action := range []struct { + name string + config map[string]any + allowAbsent bool + }{ + {"remove root", map[string]any{"remove": []any{"/previous_response_id"}}, true}, + {"remove child", map[string]any{"remove": []any{"/previous_response_id/value"}}, true}, + {"set", map[string]any{"set": map[string]any{"previous_response_id": "injected"}}, false}, + } { + t.Run(action.name, func(t *testing.T) { + rules, err := Compile([]any{action.config}) + if err != nil { + t.Fatal(err) + } + if err := rules.ValidateResponsesContinuation(); err == nil { + t.Fatal("management validation accepted a protected field override") + } + for _, value := range []string{"", "null", `""`, `"original"`} { + t.Run("value="+value, func(t *testing.T) { + body := `{"model":"gpt-5","input":"hello"` + if value != "" { + body += `,"previous_response_id":` + value + } + body += "}" + got, applied, err := rules.Apply(protocol.OpenAIResponses, execution.OperationResponsesCreate, "gpt-5", []byte(body)) + wantError := value != "" || !action.allowAbsent + if (err != nil) != wantError { + t.Fatalf("Apply error = %v, want error %t", err, wantError) + } + if err == nil && (!applied || string(got) != body) { + t.Fatalf("no-op removal changed the request: %s, applied %t", got, applied) + } + }) + } + }) + } +} diff --git a/internal/parameteroverride/rules.go b/internal/parameteroverride/rules.go index faa00b434..3dea090ef 100644 --- a/internal/parameteroverride/rules.go +++ b/internal/parameteroverride/rules.go @@ -201,7 +201,7 @@ func (rules Rules) Empty() bool { return len(rules.entries) == 0 } func (rules Rules) ValidateResponsesContinuation() error { for _, entry := range rules.entries { if entry.clientProtocol == "" || entry.clientProtocol == protocol.OpenAIResponses { - if err := entry.validateResponsesContinuation(); err != nil { + if err := entry.validateResponsesContinuation(nil); err != nil { return err } } @@ -209,7 +209,7 @@ func (rules Rules) ValidateResponsesContinuation() error { return nil } -func (entry rule) validateResponsesContinuation() error { +func (entry rule) validateResponsesContinuation(object *requestValue) error { for key := range entry.set { if strings.EqualFold(key, "previous_response_id") { return fmt.Errorf("parameter overrides cannot change previous_response_id") @@ -217,6 +217,15 @@ func (entry rule) validateResponsesContinuation() error { } for _, path := range entry.remove { if strings.EqualFold(path[0], "previous_response_id") { + if object != nil { + if err := object.load(); err != nil { + return err + } + // 旧规则删除不存在的字段不改变路由;保存配置时仍严格拒绝。 + if object.currentField(path[0]).raw == nil { + continue + } + } return fmt.Errorf("parameter overrides cannot change previous_response_id") } } @@ -251,11 +260,6 @@ func (rules Rules) Apply( matched := make([]rule, 0, len(rules.entries)) for _, entry := range rules.entries { if entry.matches(clientProtocol, clientModel) { - if clientProtocol == protocol.OpenAIResponses { - if err := entry.validateResponsesContinuation(); err != nil { - return nil, false, err - } - } matched = append(matched, entry) } } @@ -273,6 +277,13 @@ func (rules Rules) Apply( object.planPath(pointer) } } + if clientProtocol == protocol.OpenAIResponses { + for _, entry := range matched { + if err := entry.validateResponsesContinuation(object); err != nil { + return nil, false, err + } + } + } for _, entry := range matched { for _, pointer := range entry.remove { if err := object.remove(pointer); err != nil { diff --git a/internal/state/response_bindings.go b/internal/state/response_bindings.go index 7a04b065e..9cebc3737 100644 --- a/internal/state/response_bindings.go +++ b/internal/state/response_bindings.go @@ -11,6 +11,7 @@ import ( const ( DefaultResponseBindingTTL = 30 * 24 * time.Hour DefaultResponseBindingCapacity = 100_000 + maxResponseIDBytes = 4 << 10 maxResponseBindingIDBytes = 16 << 20 ) @@ -70,7 +71,7 @@ func (bindings *ResponseBindings) Lookup(accessKeyID uint, responseID string) (R // Record 在响应下发前登记;不同归属冲突时拒绝当前响应,不覆盖已有归属。 func (bindings *ResponseBindings) Record(accessKeyID uint, responseID string, ref CredentialRef) bool { if bindings == nil || accessKeyID == 0 || responseID == "" || ref.ID == 0 || - ref.GroupID == 0 || ref.IdentityGeneration == 0 || len(responseID) > maxResponseBindingIDBytes { + ref.GroupID == 0 || ref.IdentityGeneration == 0 || len(responseID) > maxResponseIDBytes { return false } bindings.mu.Lock() @@ -124,7 +125,7 @@ func (bindings *ResponseBindings) RestoreCheckpoint(checkpoint []ResponseBinding for _, binding := range ordered { if binding.AccessKeyID == 0 || binding.ResponseID == "" || binding.CredentialID == 0 || binding.GroupID == 0 || binding.IdentityGeneration == 0 || !binding.ExpiresAt.After(now) || - len(binding.ResponseID) > maxResponseBindingIDBytes { + len(binding.ResponseID) > maxResponseIDBytes { continue } if !bindings.insert(binding) { diff --git a/internal/state/response_bindings_test.go b/internal/state/response_bindings_test.go index 1b6a124c0..d75fa47e4 100644 --- a/internal/state/response_bindings_test.go +++ b/internal/state/response_bindings_test.go @@ -1,6 +1,7 @@ package state import ( + "strconv" "strings" "sync" "testing" @@ -54,15 +55,59 @@ func TestResponseBindingsEvictOldestAndBoundIDMemory(t *testing.T) { t.Fatalf("recent binding %q was lost", id) } } - largeID := strings.Repeat("x", maxResponseBindingIDBytes) - if !bindings.Record(1, largeID, ref) { - t.Fatal("opaque ID within memory budget was rejected") + bindings = NewResponseBindings() + prefix := strings.Repeat("x", 4<<10) + var firstID, lastID string + for index := range maxResponseBindingIDBytes/(4<<10) + 1 { + suffix := strconv.Itoa(index) + lastID = prefix[len(suffix):] + suffix + if index == 0 { + firstID = lastID + } + if !bindings.Record(1, lastID, ref) { + t.Fatal("valid ID was rejected before eviction") + } + } + if _, ok := bindings.Lookup(1, firstID); ok { + t.Fatal("ID byte budget retained the oldest entry") + } + if _, ok := bindings.Lookup(1, lastID); !ok { + t.Fatal("ID byte budget lost the newest entry") + } +} + +func TestResponseBindingsRejectOversizedIDWithoutEvictingOtherAccessKeys(t *testing.T) { + bindings := NewResponseBindings() + ref := CredentialRef{ID: 1, GroupID: 1, IdentityGeneration: 1} + maxID := strings.Repeat("x", 4<<10) + if !bindings.Record(1, maxID, ref) { + t.Fatal("ID at the single-record limit was rejected") + } + for _, size := range []int{len(maxID) + 1, 16 << 20} { + if bindings.Record(2, strings.Repeat("y", size), ref) { + t.Errorf("accepted oversized ID of %d bytes", size) + } + if _, ok := bindings.Lookup(1, maxID); !ok { + t.Error("oversized ID evicted another AccessKey's binding") + } + } +} + +func TestResponseBindingsCheckpointSkipsOversizedIDs(t *testing.T) { + bindings := NewResponseBindings() + valid := ResponseBinding{AccessKeyID: 1, ResponseID: strings.Repeat("x", 4<<10), + CredentialID: 1, GroupID: 1, IdentityGeneration: 1, ExpiresAt: time.Now().Add(time.Hour)} + oversized := valid + oversized.AccessKeyID = 2 + oversized.ResponseID += "x" + if err := bindings.RestoreCheckpoint([]ResponseBinding{valid, oversized}); err != nil { + t.Fatal(err) } - if _, ok := bindings.Lookup(1, "newest"); ok { - t.Fatal("ID byte budget did not evict old entries") + if _, ok := bindings.Lookup(1, valid.ResponseID); !ok { + t.Fatal("restore discarded a valid binding") } - if bindings.Record(2, largeID+"x", ref) { - t.Fatal("oversized record exceeded ID memory budget") + if _, ok := bindings.Lookup(2, oversized.ResponseID); ok { + t.Fatal("restore accepted an oversized ID") } } From a9bc90fabb421dd59573364e32fc4c38971eef86 Mon Sep 17 00:00:00 2001 From: tbphp Date: Wed, 9 Sep 2026 22:12:03 +0800 Subject: [PATCH 3/3] =?UTF-8?q?fix(responses):=20=E6=8C=89=E5=8D=8F?= =?UTF-8?q?=E8=AE=AE=E5=AD=98=E5=82=A8=E8=83=BD=E5=8A=9B=E5=8C=B9=E9=85=8D?= =?UTF-8?q?=E7=BB=AD=E6=8E=A5=E8=B7=AF=E7=94=B1?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.md | 2 +- README_CN.md | 2 +- README_JP.md | 2 +- internal/dialect/request_execution.go | 17 ++-- internal/dialect/request_execution_test.go | 1 + internal/execution/bifrost/executor.go | 21 +++-- internal/execution/contracts.go | 38 +++++---- internal/execution/contracts_test.go | 10 +++ internal/execution/validation.go | 8 ++ internal/gateway/execution_forward.go | 1 + internal/gateway/forward.go | 1 + internal/gateway/handler.go | 1 + internal/gateway/responses_continuation.go | 6 +- .../gateway/responses_continuation_test.go | 84 ++++++++++++++++++- internal/scheduler/channel_scheduler_test.go | 43 +++++++++- internal/scheduler/inspect.go | 8 ++ 16 files changed, 205 insertions(+), 40 deletions(-) diff --git a/README.md b/README.md index 00f14a977..a7e3094d5 100644 --- a/README.md +++ b/README.md @@ -239,7 +239,7 @@ Environment proxies apply only when no proxy is specified on the credential, gro - 2.0 is designed for a **single application instance**. Instances do not share state, so horizontal scaling is not supported. - Usage and cost are **estimates** derived from upstream responses. They support operational analysis and capacity planning, and do not equal a provider invoice or a financial reconciliation. - Subscription channels depend on upstream OAuth and compatibility protocols and may change as upstreams change. Only connect accounts you are entitled to use, and follow each provider's terms. -- Responses continuation with `previous_response_id` currently supports the `openai` and `gpt_load` channels. Response ownership is isolated by AccessKey and pins the original credential when current routing permits, independently of soft affinity. Unknown IDs, including IDs created before upgrading or outside this gateway, are rejected. Group parameter overrides cannot change this field. +- Responses continuation with `previous_response_id` automatically uses native Responses routes that declare upstream-managed storage: currently `openai`, `gpt_load`, `xai`, `newapi`, `cliproxyapi`, and `sub2api`. Ownership is isolated by AccessKey and pins the original credential when current routing permits, independently of soft affinity; actual state availability depends on the upstream. Stateless and converted responses are not registered, and Codex subscription WebSocket continuation is not yet integrated. Unknown IDs, including IDs created before upgrading or outside this gateway, are rejected. Group parameter overrides cannot change this field. - Response bindings stay in memory for up to 30 days, with limits of 100,000 entries and 16 MiB of ID text; older entries are evicted when capacity is reached. A successful checkpoint during normal shutdown allows restoration from the same data directory. Crash recovery and continued upstream state availability are not guaranteed. - `conversation` and other existing resource IDs are outside this ownership routing scope and still depend on a single credential or upstream resource sharing across credentials. diff --git a/README_CN.md b/README_CN.md index acadf4e7b..7ecb35432 100644 --- a/README_CN.md +++ b/README_CN.md @@ -239,7 +239,7 @@ Windows 普通用户可改为下载 `gpt-load-windows-setup.exe`。双击并确 - 2.0 按**单应用实例**设计,多个实例之间不共享状态,不支持直接横向扩容。 - 用量与成本是基于上游返回数据的**估算**,用于运行分析和资源评估,不等同于服务商账单或财务对账结果。 - 订阅渠道依赖上游 OAuth 与兼容协议,可能随上游变化调整。请只接入自己有权使用的账号,并遵守对应服务商条款。 -- Responses 的 `previous_response_id` 续接目前支持 `openai`、`gpt_load` 渠道:按 AccessKey 隔离响应归属,并在当前路由允许时固定原凭据,不受软亲和开关影响。未知 ID(包括升级前或网关外创建的 ID)直接拒绝;Group 参数覆盖不能改写该字段。 +- Responses 的 `previous_response_id` 续接按协议及现有存储能力自动接入:原生 Responses 且声明由上游管理状态的渠道目前包括 `openai`、`gpt_load`、`xai`、`newapi`、`cliproxyapi`、`sub2api`。按 AccessKey 隔离归属,在当前路由允许时固定原凭据,不受软亲和开关影响;实际状态可用性由上游决定。无状态及转换响应不登记,Codex 订阅的 WebSocket 续接尚未接入。未知 ID(包括升级前或网关外创建的 ID)直接拒绝;Group 参数覆盖不能改写该字段。 - 响应归属保存在内存中,默认保留 30 天,最多 100,000 条,ID 文本合计最多 16 MiB,达到容量时淘汰旧记录。正常停机成功保存 checkpoint 后可在同一数据目录恢复;不保证崩溃恢复或上游历史仍有效。 - `conversation` 与其他既有资源 ID 不在上述归属路由范围内,仍依赖单凭据或上游跨凭据共享资源。 diff --git a/README_JP.md b/README_JP.md index 384d65e2f..9ff0a9320 100644 --- a/README_JP.md +++ b/README_JP.md @@ -239,7 +239,7 @@ Windows の一般ユーザーは代わりに `gpt-load-windows-setup.exe` を利 - 2.0 は**単一アプリケーションインスタンス**を前提に設計されています。インスタンス間で状態を共有しないため、そのままの水平スケールには対応していません。 - 使用量とコストはアップストリームの応答に基づく**概算**です。運用分析やリソース評価には使えますが、プロバイダーの請求書や会計上の照合結果とは一致しません。 - サブスクリプションチャネルはアップストリームの OAuth と互換プロトコルに依存し、アップストリームの変更に伴って調整が必要になる場合があります。利用権限のあるアカウントのみを接続し、各プロバイダーの規約に従ってください。 -- Responses の `previous_response_id` による継続は、現在 `openai`、`gpt_load` チャネルに対応しています。応答の帰属を AccessKey ごとに分離し、現在のルーティングで許可される場合に元の認証情報へ固定します。ソフトアフィニティ設定には依存しません。アップグレード前やゲートウェイ外で作成されたものを含め、不明な ID は拒否されます。Group のパラメータ上書きでこのフィールドを変更することはできません。 +- Responses の `previous_response_id` による継続は、ネイティブ Responses と上流での状態管理を宣言した経路に自動で適用されます。現在は `openai`、`gpt_load`、`xai`、`newapi`、`cliproxyapi`、`sub2api` が該当します。帰属を AccessKey ごとに分離し、現在のルーティングで許可される元の認証情報へ固定します。ソフトアフィニティ設定には依存せず、状態が実際に利用できるかは上流に依存します。ステートレス応答と変換された応答は登録せず、Codex サブスクリプションの WebSocket 継続はまだ接続していません。アップグレード前やゲートウェイ外で作成されたものを含め、不明な ID は拒否されます。Group のパラメータ上書きでこのフィールドを変更することはできません。 - 応答の帰属はメモリに最大 30 日間保持され、上限は 100,000 件および ID テキスト合計 16 MiB です。容量に達すると古い記録を削除します。通常終了時に checkpoint の保存が成功すれば、同じデータディレクトリから復元できます。クラッシュからの復元や、アップストリームの履歴が引き続き有効であることは保証しません。 - `conversation` とその他の既存リソース ID はこの帰属ルーティングの対象外であり、単一の認証情報またはアップストリームでの認証情報間のリソース共有が引き続き必要です。 diff --git a/internal/dialect/request_execution.go b/internal/dialect/request_execution.go index e4e2173af..16f1b87f6 100644 --- a/internal/dialect/request_execution.go +++ b/internal/dialect/request_execution.go @@ -115,8 +115,7 @@ func responsesCreateRequirements( if !ok { return execution.RouteRequirementNative, execution.ResponsesStorePreferenceNone } - if hasMeaningfulField(root, "previous_response_id") || - hasMeaningfulField(root, "conversation") || + if hasMeaningfulField(root, "conversation") || responsesPromptReferencesProviderResource(root["prompt"]) { return execution.RouteRequirementNative, execution.ResponsesStorePreferenceNone } @@ -128,17 +127,15 @@ func responsesCreateRequirements( return execution.RouteRequirementNative, execution.ResponsesStorePreferenceNone } value, exists := root["store"] - if !exists { - return execution.RouteRequirementAny, execution.ResponsesStorePreferencePreferStored - } - if value == nil { - return execution.RouteRequirementNative, execution.ResponsesStorePreferenceNone - } store, ok := value.(bool) - if !ok { + if exists && !ok { return execution.RouteRequirementNative, execution.ResponsesStorePreferenceNone } - if store { + // 其他资源与参数约束优先;仅纯 ID 续接使用上游存储要求。 + if hasMeaningfulField(root, "previous_response_id") { + return execution.RouteRequirementNative, execution.ResponsesStorePreferenceRequireStored + } + if !exists || store { return execution.RouteRequirementAny, execution.ResponsesStorePreferencePreferStored } return execution.RouteRequirementAny, execution.ResponsesStorePreferenceNone diff --git a/internal/dialect/request_execution_test.go b/internal/dialect/request_execution_test.go index e15e90cfa..c1625672d 100644 --- a/internal/dialect/request_execution_test.go +++ b/internal/dialect/request_execution_test.go @@ -405,6 +405,7 @@ func TestResponsesCreateClassifiesStorePreferenceWithoutWeakeningNativeResources name: "previous response remains native only", body: `{"model":"gpt-5","previous_response_id":"resp_123","store":false}`, routeRequirement: execution.RouteRequirementNative, + storePreference: execution.ResponsesStorePreferenceRequireStored, }, { name: "provider resource outranks omitted store", diff --git a/internal/execution/bifrost/executor.go b/internal/execution/bifrost/executor.go index 4fcabaa50..c2ac25e75 100644 --- a/internal/execution/bifrost/executor.go +++ b/internal/execution/bifrost/executor.go @@ -427,13 +427,20 @@ func (r *Runtime) prepare(spec execution.AttemptSpec, stream bool) (preparedAtte return preparedAttempt{}, &failure } if spec.Operation == execution.OperationResponsesCreate && - spec.RouteRequirement.Normalize() == execution.RouteRequirementNative && - !resolved.SupportsResponsesLifecycle() { - failure := notSentConversionFailure( - execution.ErrorCodeCriticalSemanticLoss, - "target does not support Responses resource lifecycle", - ) - return preparedAttempt{}, &failure + spec.RouteRequirement.Normalize() == execution.RouteRequirementNative { + var stateSupported bool + if spec.ResponsesStorePreference == execution.ResponsesStorePreferenceRequireStored { + stateSupported = resolved.ResponsesStoreHandling(spec.ClientProtocol, spec.Operation) == channel.ResponsesStoreHandlingUpstreamManaged + } else { + stateSupported = resolved.SupportsResponsesLifecycle() + } + if !stateSupported { + failure := notSentConversionFailure( + execution.ErrorCodeCriticalSemanticLoss, + "target does not support required Responses state", + ) + return preparedAttempt{}, &failure + } } if spec.Operation == execution.OperationResponsesPassthrough && providerKind != channel.ProviderOpenAI && diff --git a/internal/execution/contracts.go b/internal/execution/contracts.go index acc57e484..4d42df5c6 100644 --- a/internal/execution/contracts.go +++ b/internal/execution/contracts.go @@ -148,18 +148,21 @@ func (r RouteRequirement) Allows(mode RouteMode) bool { return r.Normalize() == RouteRequirementAny || mode == RouteNative } -// ResponsesStorePreference records the one Responses Create storage intent -// that may use an explicitly declared stateless compatibility route. +// ResponsesStorePreference records the Responses Create storage requirement +// and whether an explicitly declared stateless compatibility route is allowed. type ResponsesStorePreference string const ( ResponsesStorePreferenceNone ResponsesStorePreference = "" ResponsesStorePreferencePreferStored ResponsesStorePreference = "prefer_stored" + // 续接上一轮响应必须保留状态语义,不允许降级为无状态。 + ResponsesStorePreferenceRequireStored ResponsesStorePreference = "require_stored" ) // Valid reports whether the Responses storage preference is recognized. func (p ResponsesStorePreference) Valid() bool { - return p == ResponsesStorePreferenceNone || p == ResponsesStorePreferencePreferStored + return p == ResponsesStorePreferenceNone || p == ResponsesStorePreferencePreferStored || + p == ResponsesStorePreferenceRequireStored } // CredentialSnapshot is the exact logical credential selected for an attempt. @@ -218,20 +221,21 @@ type AttemptTimeouts struct { // NewAttemptSpec or Clone must be used at ownership boundaries because Query, // Header, Body, TargetConfig, and Credential contain reference-backed values. type AttemptSpec struct { - RequestID string `json:"request_id"` - AttemptID string `json:"attempt_id"` - Sequence uint32 `json:"sequence"` - ChannelID string `json:"channel_id"` - RouteMode RouteMode `json:"route_mode"` - ClientProtocol protocol.Protocol `json:"client_protocol"` - Operation Operation `json:"operation"` - RouteRequirement RouteRequirement `json:"route_requirement"` - ResponsesStoreDowngraded bool `json:"responses_store_downgraded,omitempty"` - ClientModel string `json:"client_model,omitempty"` - UpstreamModel string `json:"upstream_model,omitempty"` - Method string `json:"method"` - Path string `json:"path"` - Query url.Values `json:"query,omitempty"` + RequestID string `json:"request_id"` + AttemptID string `json:"attempt_id"` + Sequence uint32 `json:"sequence"` + ChannelID string `json:"channel_id"` + RouteMode RouteMode `json:"route_mode"` + ClientProtocol protocol.Protocol `json:"client_protocol"` + Operation Operation `json:"operation"` + RouteRequirement RouteRequirement `json:"route_requirement"` + ResponsesStorePreference ResponsesStorePreference `json:"responses_store_preference,omitempty"` + ResponsesStoreDowngraded bool `json:"responses_store_downgraded,omitempty"` + ClientModel string `json:"client_model,omitempty"` + UpstreamModel string `json:"upstream_model,omitempty"` + Method string `json:"method"` + Path string `json:"path"` + Query url.Values `json:"query,omitempty"` // RawQuery preserves the original query bytes when exact forwarding matters. // It is mutually exclusive with Query and intentionally is not URL-decoded. RawQuery string `json:"raw_query,omitempty"` diff --git a/internal/execution/contracts_test.go b/internal/execution/contracts_test.go index aeb3d7b97..e89d13d5b 100644 --- a/internal/execution/contracts_test.go +++ b/internal/execution/contracts_test.go @@ -78,6 +78,7 @@ func TestOperationAndDispatchEnums(t *testing.T) { for _, preference := range []ResponsesStorePreference{ ResponsesStorePreferenceNone, ResponsesStorePreferencePreferStored, + ResponsesStorePreferenceRequireStored, } { if !preference.Valid() { t.Fatalf("expected Responses store preference %q to be valid", preference) @@ -402,6 +403,15 @@ func TestValidationAcceptsValidContractsAndRejectsInvalidFields(t *testing.T) { {name: "request id", mutate: func(s *AttemptSpec) { s.RequestID = "" }, field: "request_id"}, {name: "attempt id", mutate: func(s *AttemptSpec) { s.AttemptID = "" }, field: "attempt_id"}, {name: "sequence", mutate: func(s *AttemptSpec) { s.Sequence = 0 }, field: "sequence"}, + {name: "storage preference", mutate: func(s *AttemptSpec) { + s.ResponsesStorePreference = ResponsesStorePreference("invalid") + }, field: "responses_store_preference"}, + {name: "stored continuation on another protocol", mutate: func(s *AttemptSpec) { + s.ResponsesStorePreference = ResponsesStorePreferenceRequireStored + s.RouteRequirement = RouteRequirementNative + s.ClientProtocol = protocol.OpenAICompletions + s.Operation = OperationChatCompletion + }, field: "responses_store_preference"}, {name: "channel", mutate: func(s *AttemptSpec) { s.ChannelID = "" }, field: "channel_id"}, {name: "route mode", mutate: func(s *AttemptSpec) { s.RouteMode = RouteMode("fallback") }, field: "route_mode"}, {name: "route requirement", mutate: func(s *AttemptSpec) { s.RouteRequirement = RouteRequirement("converted-only") }, field: "route_requirement"}, diff --git a/internal/execution/validation.go b/internal/execution/validation.go index 43fa79637..670679798 100644 --- a/internal/execution/validation.go +++ b/internal/execution/validation.go @@ -68,6 +68,14 @@ func (s AttemptSpec) Validate() error { if !s.RouteRequirement.Valid() { return validationError("route_requirement", "unsupported value") } + if !s.ResponsesStorePreference.Valid() { + return validationError("responses_store_preference", "unsupported value") + } + if s.ResponsesStorePreference == ResponsesStorePreferenceRequireStored && + (s.ClientProtocol != protocol.OpenAIResponses || s.Operation != OperationResponsesCreate || + s.RouteRequirement.Normalize() != RouteRequirementNative) { + return validationError("responses_store_preference", "requires a native Responses Create request") + } if s.ResponsesStoreDowngraded && (s.ClientProtocol != protocol.OpenAIResponses || s.Operation != OperationResponsesCreate || diff --git a/internal/gateway/execution_forward.go b/internal/gateway/execution_forward.go index fb61f234a..4909dd8ae 100644 --- a/internal/gateway/execution_forward.go +++ b/internal/gateway/execution_forward.go @@ -739,6 +739,7 @@ func newExecutionAttemptSpec(input ForwardInput) (execution.AttemptSpec, error) ClientProtocol: input.ClientProtocol, Operation: input.Operation, RouteRequirement: input.RouteRequirement, + ResponsesStorePreference: input.ResponsesStorePreference, ResponsesStoreDowngraded: input.ResponsesStoreDowngraded, ClientModel: input.ExternalModel, UpstreamModel: input.UpstreamModelID, diff --git a/internal/gateway/forward.go b/internal/gateway/forward.go index 78ee94fe6..958d1c5ce 100644 --- a/internal/gateway/forward.go +++ b/internal/gateway/forward.go @@ -45,6 +45,7 @@ type ForwardInput struct { ClientProtocol protocol.Protocol Operation execution.Operation RouteRequirement execution.RouteRequirement + ResponsesStorePreference execution.ResponsesStorePreference ResponsesStoreDowngraded bool ChannelID string RouteMode execution.RouteMode diff --git a/internal/gateway/handler.go b/internal/gateway/handler.go index 44262888d..2c3674957 100644 --- a/internal/gateway/handler.go +++ b/internal/gateway/handler.go @@ -1156,6 +1156,7 @@ func (handler *Handler) executeAttempts( ClientProtocol: selectedDialect.Protocol(), Operation: originalMetadata.Operation, RouteRequirement: originalMetadata.RouteRequirement, + ResponsesStorePreference: originalMetadata.ResponsesStorePreference, ResponsesStoreDowngraded: selection.ResponsesStoreDowngraded, ChannelID: string(selection.ChannelID), RouteMode: execution.RouteMode(selection.RouteMode), diff --git a/internal/gateway/responses_continuation.go b/internal/gateway/responses_continuation.go index 6c14e6e0c..69c81b95c 100644 --- a/internal/gateway/responses_continuation.go +++ b/internal/gateway/responses_continuation.go @@ -5,8 +5,10 @@ import ( "fmt" "net/http" + "gpt-load/internal/channel" "gpt-load/internal/dialect" "gpt-load/internal/execution" + "gpt-load/internal/protocol" "gpt-load/internal/scheduler" "gpt-load/internal/state" ) @@ -19,7 +21,9 @@ func (handler *Handler) responseBindingObserver( ) func([]byte) error { if request == nil || request.Method != http.MethodPost || request.Path != "/v1/responses" || selection.RouteMode != execution.RouteNative || selection.ResponsesStoreDowngraded || - !selection.ResolvedTarget.SupportsResponsesLifecycle() { + selection.ResolvedTarget.ResponsesStoreHandling( + protocol.OpenAIResponses, execution.OperationResponsesCreate, + ) != channel.ResponsesStoreHandlingUpstreamManaged { return nil } var options struct { diff --git a/internal/gateway/responses_continuation_test.go b/internal/gateway/responses_continuation_test.go index a3e222495..07f94a0f5 100644 --- a/internal/gateway/responses_continuation_test.go +++ b/internal/gateway/responses_continuation_test.go @@ -4,15 +4,18 @@ import ( "bytes" "compress/gzip" "context" + "encoding/json" "fmt" "net/http" "net/http/httptest" + "sync/atomic" "testing" "time" "github.com/gin-gonic/gin" "gpt-load/internal/app" + "gpt-load/internal/channel" "gpt-load/internal/dialect" "gpt-load/internal/execution" "gpt-load/internal/parameteroverride" @@ -44,6 +47,54 @@ func TestResponsesContinuationPinsCredentialWithoutSoftAffinity(t *testing.T) { } } +func TestResponsesContinuationUsesNativeStorageCapabilities(t *testing.T) { + for _, channelID := range []channel.ID{ + channel.OpenAI, channel.GPTLoad, channel.XAI, channel.NewAPI, channel.CLIProxyAPI, channel.Sub2API, + } { + t.Run(string(channelID), func(t *testing.T) { + type observedRequest struct { + PreviousResponseID string `json:"previous_response_id"` + Store *bool `json:"store"` + authorization string + } + requests := make(chan observedRequest, 8) + var count atomic.Int32 + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + var observed observedRequest + if err := json.NewDecoder(request.Body).Decode(&observed); err != nil { + t.Error(err) + writer.WriteHeader(http.StatusBadRequest) + return + } + observed.authorization = request.Header.Get("Authorization") + requests <- observed + writer.Header().Set("Content-Type", "application/json") + _, _ = writer.Write(storedResponse(fmt.Sprintf("response-%d", count.Add(1))).Body) + })) + defer server.Close() + handler, engine, sink := newContinuationFixture(t, newTestExecutionForwarder(t)) + setContinuationChannel(t, handler, channelID, server.URL) + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","input":"initial"}`, http.StatusOK) + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","previous_response_id":"response-1","input":"continue"}`, http.StatusOK) + serveContinuation(t, engine, "gl-client", `{"model":"gpt-4o","previous_response_id":"response-2","input":"continue","store":false}`, http.StatusOK) + if count.Load() != 3 { + t.Fatalf("upstream attempts = %d, want 3", count.Load()) + } + for index, id := range []string{"", "response-1", "response-2"} { + observed := <-requests + if observed.PreviousResponseID != id || observed.authorization != "Bearer sk-one" || + (index == 2 && (observed.Store == nil || *observed.Store)) { + t.Fatalf("request %d lost continuation semantics: %+v", index, observed) + } + } + assertAffinityHits(t, sink.snapshot(), []bool{false, true, true}) + if _, ok := handler.responseBindings.Lookup(1, "response-3"); ok { + t.Fatal("store:false registered a new continuation ID") + } + }) + } +} + func TestResponsesContinuationRejectsUnknownAndOtherAccessKeyIDs(t *testing.T) { forwarder := &scriptedForwarder{results: []UpstreamResult{storedResponse("first")}} handler, engine, sink := newContinuationFixture(t, forwarder) @@ -116,7 +167,9 @@ func TestResponsesContinuationRegistersSSEBeforeDelivery(t *testing.T) { return execution.StreamResult{StatusCode: http.StatusOK, DispatchState: execution.DispatchMaybeSent, ResponseStarted: true} }, } - _, engine, _ = newContinuationFixture(t, NewExecutionForwarder(executor)) + handler, runtime, _ := newContinuationFixture(t, NewExecutionForwarder(executor)) + engine = runtime + setContinuationChannel(t, handler, channel.NewAPI, "https://upstream.example") request := httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewBufferString(`{"model":"gpt-4o","input":"initial","stream":true}`)) request.Header.Set("Authorization", "Bearer gl-client") engine.ServeHTTP(writer, request) @@ -293,6 +346,35 @@ func newContinuationFixture(t *testing.T, forwarder AttemptForwarder) (*Handler, return handler, engine, sink } +func setContinuationChannel(t *testing.T, handler *Handler, channelID channel.ID, baseURL string) { + t.Helper() + credentials := make([]state.CredentialConfig, 0, 2) + for _, id := range []uint{1, 2} { + credentials = append(credentials, state.CredentialConfig{ + ID: id, GroupID: 1, Status: state.CredentialStatusActive, + Version: 1, IdentityGeneration: uint64(id), Fingerprint: fmt.Sprintf("credential-%d", id), + }) + } + params, err := json.Marshal(map[string]string{"base_url": baseURL}) + if err != nil { + t.Fatal(err) + } + if _, err := handler.manager.Publish(state.CompileInput{ + ChannelRegistry: channel.NewRegistry(), + Groups: []state.GroupConfig{{ + ID: 1, Name: string(channelID), ChannelID: channelID, ConnectionType: "api_key", + Params: params, + Models: []state.ModelConfig{{ID: "gpt-4o"}}, Enabled: true, + }}, + Credentials: credentials, + AccessKeys: []state.AccessKeyConfig{{ + ID: 1, Name: "client", KeyHash: handler.encryption.Hash("gl-client"), Status: state.AccessKeyStatusActive, + }}, + }); err != nil { + t.Fatal(err) + } +} + func serveContinuation(t *testing.T, engine http.Handler, key, body string, status int) *httptest.ResponseRecorder { t.Helper() request := httptest.NewRequest(http.MethodPost, "/v1/responses", bytes.NewBufferString(body)) diff --git a/internal/scheduler/channel_scheduler_test.go b/internal/scheduler/channel_scheduler_test.go index 21f86e3a7..0b0821bd0 100644 --- a/internal/scheduler/channel_scheduler_test.go +++ b/internal/scheduler/channel_scheduler_test.go @@ -286,7 +286,7 @@ func TestRouteRequirementKeepsStatefulResponsesOnNativeTargets(t *testing.T) { } } -func TestStatefulResponsesCreateRequiresLifecycleTargetEvenWhenWireIsNative(t *testing.T) { +func TestResponsesContinuationSeparatesStorageFromOtherResourceRequirements(t *testing.T) { t.Parallel() snapshot, err := state.Compile(state.CompileInput{ @@ -304,6 +304,14 @@ func TestStatefulResponsesCreateRequiresLifecycleTargetEvenWhenWireIsNative(t *t Params: json.RawMessage(`{}`), Enabled: true, Models: []state.ModelConfig{{ID: "grok", Alias: "gpt"}}, }, + {ConnectionType: "subscription", ID: 10, Name: "codex", ChannelID: channel.Codex, + Params: json.RawMessage(`{}`), Enabled: true, + Models: []state.ModelConfig{{ID: "gpt-codex", Alias: "gpt"}}, + }, + {ConnectionType: "api_key", ID: 11, Name: "converted", ChannelID: channel.Anthropic, + Params: json.RawMessage(`{}`), Enabled: true, + Models: []state.ModelConfig{{ID: "claude", Alias: "gpt"}}, + }, }, }) if err != nil { @@ -320,6 +328,39 @@ func TestStatefulResponsesCreateRequiresLifecycleTargetEvenWhenWireIsNative(t *t if !slices.Equal(got, []uint{7}) { t.Fatalf("CandidateGroupIDsForQuery() = %#v, want lifecycle-capable OpenAI group [7]", got) } + for _, test := range []struct { + name string + fields string + want []uint + }{ + {"continuation", `,"input":"continue"`, []uint{7, 9}}, + {"unstored next response", `,"store":false`, []uint{7, 9}}, + {"conversation", `,"conversation":"conv_1"`, []uint{7}}, + {"background", `,"background":true`, []uint{7}}, + {"stored prompt", `,"prompt":{"id":"pmpt_1"}`, []uint{7}}, + {"input reference", `,"input":[{"type":"item_reference","id":"item_1"}]`, []uint{7}}, + {"file search", `,"tools":[{"type":"file_search","vector_store_ids":["vs_1"]}]`, []uint{7}}, + {"null store", `,"store":null`, []uint{7}}, + {"invalid store", `,"store":"invalid"`, []uint{7}}, + } { + t.Run(test.name, func(t *testing.T) { + metadata, err := dialect.NewOpenAIResponses().InspectRequest(&dialect.ParsedRequest{ + Method: http.MethodPost, Path: "/v1/responses", + Body: []byte(`{"model":"gpt","previous_response_id":"resp_1"` + test.fields + `}`), + }) + if err != nil { + t.Fatal(err) + } + query := Query{ + ClientProtocol: protocol.OpenAIResponses, Operation: metadata.Operation, + RouteRequirement: metadata.RouteRequirement, ResponsesStorePreference: metadata.ResponsesStorePreference, + ExternalModel: metadata.Model, + } + if got := CandidateGroupIDsForQuery(snapshot, query); !slices.Equal(got, test.want) { + t.Fatalf("candidate groups = %v, want %v", got, test.want) + } + }) + } } func TestResponsesStorePreferenceKeepsTheWholeRequestOnExactTargets(t *testing.T) { diff --git a/internal/scheduler/inspect.go b/internal/scheduler/inspect.go index d781df745..850358ff5 100644 --- a/internal/scheduler/inspect.go +++ b/internal/scheduler/inspect.go @@ -203,6 +203,14 @@ func routeRequirementSatisfied( if query.operation != execution.OperationResponsesCreate { return true, false, "" } + if query.responsesStorePreference == execution.ResponsesStorePreferenceRequireStored { + if route.Mode == channel.RouteNative && route.ResolvedTarget.ResponsesStoreHandling( + protocol.OpenAIResponses, execution.OperationResponsesCreate, + ) == channel.ResponsesStoreHandlingUpstreamManaged { + return true, false, "" + } + return false, false, ReasonNativeRouteRequired + } if query.routeRequirement.Normalize() == execution.RouteRequirementNative { if route.ResolvedTarget.SupportsResponsesLifecycle() { return true, false, ""