From e6eb23eaac8f76f60c7fbe2d707416eb4ea4e343 Mon Sep 17 00:00:00 2001 From: song Date: Sat, 25 Jul 2026 07:39:32 +0800 Subject: [PATCH] feat(openai): add Live gateway support --- backend/ent/group.go | 13 +- backend/ent/group/group.go | 10 + backend/ent/group/where.go | 15 + backend/ent/group_create.go | 65 ++ backend/ent/group_update.go | 34 + backend/ent/migrate/schema.go | 1 + backend/ent/mutation.go | 56 +- backend/ent/runtime/runtime.go | 20 +- backend/ent/schema/group.go | 3 + backend/internal/config/config.go | 10 + .../internal/handler/admin/group_handler.go | 4 + backend/internal/handler/dto/mappers.go | 1 + backend/internal/handler/dto/types.go | 2 + backend/internal/handler/openai_live.go | 231 ++++++ backend/internal/handler/openai_live_test.go | 104 +++ backend/internal/repository/api_key_repo.go | 2 + .../internal/repository/concurrency_cache.go | 165 +++- .../concurrency_cache_integration_test.go | 36 + .../repository/concurrency_cache_live_test.go | 71 ++ backend/internal/repository/gateway_cache.go | 142 ++++ .../repository/gateway_cache_live_test.go | 58 ++ backend/internal/repository/group_repo.go | 2 + .../migrations_schema_integration_test.go | 3 + backend/internal/server/api_contract_test.go | 1 + backend/internal/server/routes/gateway.go | 4 + .../prompt_audit_route_coverage_test.go | 2 + backend/internal/service/account.go | 6 + backend/internal/service/admin_group.go | 10 + .../internal/service/admin_group_duplicate.go | 1 + .../service/admin_group_duplicate_test.go | 1 + backend/internal/service/admin_service.go | 2 + .../service/admin_service_group_test.go | 4 + .../internal/service/api_key_auth_cache.go | 1 + .../service/api_key_auth_cache_impl.go | 4 +- backend/internal/service/group.go | 1 + backend/internal/service/openai_live.go | 747 ++++++++++++++++++ .../service/openai_live_lifecycle_test.go | 408 ++++++++++ backend/internal/service/openai_live_test.go | 181 +++++ backend/internal/service/openai_live_types.go | 90 +++ backend/internal/service/usage_log.go | 9 +- .../186_allow_live_usage_request_type.sql | 6 + .../migrations/187_add_group_allow_live.sql | 1 + deploy/config.example.yaml | 3 + .../components/admin/usage/UsageFilters.vue | 1 + .../src/components/admin/usage/UsageTable.vue | 2 + .../src/i18n/locales/en/admin/overview.ts | 5 + frontend/src/i18n/locales/en/dashboard.ts | 1 + .../src/i18n/locales/zh/admin/overview.ts | 5 + frontend/src/i18n/locales/zh/dashboard.ts | 1 + frontend/src/types/index.ts | 6 +- frontend/src/utils/errorBadges.ts | 5 +- frontend/src/utils/usageRequestType.ts | 4 +- frontend/src/views/admin/GroupsView.vue | 74 ++ frontend/src/views/admin/UsageView.vue | 1 + frontend/src/views/user/UsageView.vue | 2 + 55 files changed, 2607 insertions(+), 30 deletions(-) create mode 100644 backend/internal/handler/openai_live.go create mode 100644 backend/internal/handler/openai_live_test.go create mode 100644 backend/internal/repository/concurrency_cache_live_test.go create mode 100644 backend/internal/repository/gateway_cache_live_test.go create mode 100644 backend/internal/service/openai_live.go create mode 100644 backend/internal/service/openai_live_lifecycle_test.go create mode 100644 backend/internal/service/openai_live_test.go create mode 100644 backend/internal/service/openai_live_types.go create mode 100644 backend/migrations/186_allow_live_usage_request_type.sql create mode 100644 backend/migrations/187_add_group_allow_live.sql diff --git a/backend/ent/group.go b/backend/ent/group.go index 4cff508698..d915740125 100644 --- a/backend/ent/group.go +++ b/backend/ent/group.go @@ -105,6 +105,8 @@ type Group struct { SortOrder int `json:"sort_order,omitempty"` // 是否允许 /v1/messages 调度到此 OpenAI 分组 AllowMessagesDispatch bool `json:"allow_messages_dispatch,omitempty"` + // 是否允许此 OpenAI 分组访问 Live 接口 + AllowLive bool `json:"allow_live,omitempty"` // 仅允许非 apikey 类型账号关联到此分组 RequireOauthOnly bool `json:"require_oauth_only,omitempty"` // 调度时仅允许 privacy 已成功设置的账号 @@ -229,7 +231,7 @@ func (*Group) scanValues(columns []string) ([]any, error) { switch columns[i] { case group.FieldModelRouting, group.FieldSupportedModelScopes, group.FieldMessagesDispatchModelConfig, group.FieldModelsListConfig, group.FieldReasoningEffortMappings: values[i] = new([]byte) - case group.FieldPeakRateEnabled, group.FieldIsExclusive, group.FieldAllowImageGeneration, group.FieldAllowBatchImageGeneration, group.FieldImageRateIndependent, group.FieldVideoRateIndependent, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch, group.FieldRequireOauthOnly, group.FieldRequirePrivacySet: + case group.FieldPeakRateEnabled, group.FieldIsExclusive, group.FieldAllowImageGeneration, group.FieldAllowBatchImageGeneration, group.FieldImageRateIndependent, group.FieldVideoRateIndependent, group.FieldClaudeCodeOnly, group.FieldModelRoutingEnabled, group.FieldMcpXMLInject, group.FieldAllowMessagesDispatch, group.FieldAllowLive, group.FieldRequireOauthOnly, group.FieldRequirePrivacySet: values[i] = new(sql.NullBool) case group.FieldRateMultiplier, group.FieldPeakRateMultiplier, group.FieldDailyLimitUsd, group.FieldWeeklyLimitUsd, group.FieldMonthlyLimitUsd, group.FieldImageRateMultiplier, group.FieldImagePrice1k, group.FieldImagePrice2k, group.FieldImagePrice4k, group.FieldBatchImageDiscountMultiplier, group.FieldBatchImageHoldMultiplier, group.FieldVideoRateMultiplier, group.FieldVideoPrice480p, group.FieldVideoPrice720p, group.FieldVideoPrice1080p, group.FieldWebSearchPricePerCall: values[i] = new(sql.NullFloat64) @@ -537,6 +539,12 @@ func (_m *Group) assignValues(columns []string, values []any) error { } else if value.Valid { _m.AllowMessagesDispatch = value.Bool } + case group.FieldAllowLive: + if value, ok := values[i].(*sql.NullBool); !ok { + return fmt.Errorf("unexpected type %T for field allow_live", values[i]) + } else if value.Valid { + _m.AllowLive = value.Bool + } case group.FieldRequireOauthOnly: if value, ok := values[i].(*sql.NullBool); !ok { return fmt.Errorf("unexpected type %T for field require_oauth_only", values[i]) @@ -826,6 +834,9 @@ func (_m *Group) String() string { builder.WriteString("allow_messages_dispatch=") builder.WriteString(fmt.Sprintf("%v", _m.AllowMessagesDispatch)) builder.WriteString(", ") + builder.WriteString("allow_live=") + builder.WriteString(fmt.Sprintf("%v", _m.AllowLive)) + builder.WriteString(", ") builder.WriteString("require_oauth_only=") builder.WriteString(fmt.Sprintf("%v", _m.RequireOauthOnly)) builder.WriteString(", ") diff --git a/backend/ent/group/group.go b/backend/ent/group/group.go index cdf0bb9cd3..6ec01aed53 100644 --- a/backend/ent/group/group.go +++ b/backend/ent/group/group.go @@ -102,6 +102,8 @@ const ( FieldSortOrder = "sort_order" // FieldAllowMessagesDispatch holds the string denoting the allow_messages_dispatch field in the database. FieldAllowMessagesDispatch = "allow_messages_dispatch" + // FieldAllowLive holds the string denoting the allow_live field in the database. + FieldAllowLive = "allow_live" // FieldRequireOauthOnly holds the string denoting the require_oauth_only field in the database. FieldRequireOauthOnly = "require_oauth_only" // FieldRequirePrivacySet holds the string denoting the require_privacy_set field in the database. @@ -236,6 +238,7 @@ var Columns = []string{ FieldSupportedModelScopes, FieldSortOrder, FieldAllowMessagesDispatch, + FieldAllowLive, FieldRequireOauthOnly, FieldRequirePrivacySet, FieldDefaultMappedModel, @@ -341,6 +344,8 @@ var ( DefaultSortOrder int // DefaultAllowMessagesDispatch holds the default value on creation for the "allow_messages_dispatch" field. DefaultAllowMessagesDispatch bool + // DefaultAllowLive holds the default value on creation for the "allow_live" field. + DefaultAllowLive bool // DefaultRequireOauthOnly holds the default value on creation for the "require_oauth_only" field. DefaultRequireOauthOnly bool // DefaultRequirePrivacySet holds the default value on creation for the "require_privacy_set" field. @@ -576,6 +581,11 @@ func ByAllowMessagesDispatch(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldAllowMessagesDispatch, opts...).ToFunc() } +// ByAllowLive orders the results by the allow_live field. +func ByAllowLive(opts ...sql.OrderTermOption) OrderOption { + return sql.OrderByField(FieldAllowLive, opts...).ToFunc() +} + // ByRequireOauthOnly orders the results by the require_oauth_only field. func ByRequireOauthOnly(opts ...sql.OrderTermOption) OrderOption { return sql.OrderByField(FieldRequireOauthOnly, opts...).ToFunc() diff --git a/backend/ent/group/where.go b/backend/ent/group/where.go index 9320209433..4ad100573c 100644 --- a/backend/ent/group/where.go +++ b/backend/ent/group/where.go @@ -260,6 +260,11 @@ func AllowMessagesDispatch(v bool) predicate.Group { return predicate.Group(sql.FieldEQ(FieldAllowMessagesDispatch, v)) } +// AllowLive applies equality check predicate on the "allow_live" field. It's identical to AllowLiveEQ. +func AllowLive(v bool) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldAllowLive, v)) +} + // RequireOauthOnly applies equality check predicate on the "require_oauth_only" field. It's identical to RequireOauthOnlyEQ. func RequireOauthOnly(v bool) predicate.Group { return predicate.Group(sql.FieldEQ(FieldRequireOauthOnly, v)) @@ -1985,6 +1990,16 @@ func AllowMessagesDispatchNEQ(v bool) predicate.Group { return predicate.Group(sql.FieldNEQ(FieldAllowMessagesDispatch, v)) } +// AllowLiveEQ applies the EQ predicate on the "allow_live" field. +func AllowLiveEQ(v bool) predicate.Group { + return predicate.Group(sql.FieldEQ(FieldAllowLive, v)) +} + +// AllowLiveNEQ applies the NEQ predicate on the "allow_live" field. +func AllowLiveNEQ(v bool) predicate.Group { + return predicate.Group(sql.FieldNEQ(FieldAllowLive, v)) +} + // RequireOauthOnlyEQ applies the EQ predicate on the "require_oauth_only" field. func RequireOauthOnlyEQ(v bool) predicate.Group { return predicate.Group(sql.FieldEQ(FieldRequireOauthOnly, v)) diff --git a/backend/ent/group_create.go b/backend/ent/group_create.go index bfd104bcee..d12efb838a 100644 --- a/backend/ent/group_create.go +++ b/backend/ent/group_create.go @@ -607,6 +607,20 @@ func (_c *GroupCreate) SetNillableAllowMessagesDispatch(v *bool) *GroupCreate { return _c } +// SetAllowLive sets the "allow_live" field. +func (_c *GroupCreate) SetAllowLive(v bool) *GroupCreate { + _c.mutation.SetAllowLive(v) + return _c +} + +// SetNillableAllowLive sets the "allow_live" field if the given value is not nil. +func (_c *GroupCreate) SetNillableAllowLive(v *bool) *GroupCreate { + if v != nil { + _c.SetAllowLive(*v) + } + return _c +} + // SetRequireOauthOnly sets the "require_oauth_only" field. func (_c *GroupCreate) SetRequireOauthOnly(v bool) *GroupCreate { _c.mutation.SetRequireOauthOnly(v) @@ -948,6 +962,10 @@ func (_c *GroupCreate) defaults() error { v := group.DefaultAllowMessagesDispatch _c.mutation.SetAllowMessagesDispatch(v) } + if _, ok := _c.mutation.AllowLive(); !ok { + v := group.DefaultAllowLive + _c.mutation.SetAllowLive(v) + } if _, ok := _c.mutation.RequireOauthOnly(); !ok { v := group.DefaultRequireOauthOnly _c.mutation.SetRequireOauthOnly(v) @@ -1101,6 +1119,9 @@ func (_c *GroupCreate) check() error { if _, ok := _c.mutation.AllowMessagesDispatch(); !ok { return &ValidationError{Name: "allow_messages_dispatch", err: errors.New(`ent: missing required field "Group.allow_messages_dispatch"`)} } + if _, ok := _c.mutation.AllowLive(); !ok { + return &ValidationError{Name: "allow_live", err: errors.New(`ent: missing required field "Group.allow_live"`)} + } if _, ok := _c.mutation.RequireOauthOnly(); !ok { return &ValidationError{Name: "require_oauth_only", err: errors.New(`ent: missing required field "Group.require_oauth_only"`)} } @@ -1334,6 +1355,10 @@ func (_c *GroupCreate) createSpec() (*Group, *sqlgraph.CreateSpec) { _spec.SetField(group.FieldAllowMessagesDispatch, field.TypeBool, value) _node.AllowMessagesDispatch = value } + if value, ok := _c.mutation.AllowLive(); ok { + _spec.SetField(group.FieldAllowLive, field.TypeBool, value) + _node.AllowLive = value + } if value, ok := _c.mutation.RequireOauthOnly(); ok { _spec.SetField(group.FieldRequireOauthOnly, field.TypeBool, value) _node.RequireOauthOnly = value @@ -2224,6 +2249,18 @@ func (u *GroupUpsert) UpdateAllowMessagesDispatch() *GroupUpsert { return u } +// SetAllowLive sets the "allow_live" field. +func (u *GroupUpsert) SetAllowLive(v bool) *GroupUpsert { + u.Set(group.FieldAllowLive, v) + return u +} + +// UpdateAllowLive sets the "allow_live" field to the value that was provided on create. +func (u *GroupUpsert) UpdateAllowLive() *GroupUpsert { + u.SetExcluded(group.FieldAllowLive) + return u +} + // SetRequireOauthOnly sets the "require_oauth_only" field. func (u *GroupUpsert) SetRequireOauthOnly(v bool) *GroupUpsert { u.Set(group.FieldRequireOauthOnly, v) @@ -3193,6 +3230,20 @@ func (u *GroupUpsertOne) UpdateAllowMessagesDispatch() *GroupUpsertOne { }) } +// SetAllowLive sets the "allow_live" field. +func (u *GroupUpsertOne) SetAllowLive(v bool) *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.SetAllowLive(v) + }) +} + +// UpdateAllowLive sets the "allow_live" field to the value that was provided on create. +func (u *GroupUpsertOne) UpdateAllowLive() *GroupUpsertOne { + return u.Update(func(s *GroupUpsert) { + s.UpdateAllowLive() + }) +} + // SetRequireOauthOnly sets the "require_oauth_only" field. func (u *GroupUpsertOne) SetRequireOauthOnly(v bool) *GroupUpsertOne { return u.Update(func(s *GroupUpsert) { @@ -4345,6 +4396,20 @@ func (u *GroupUpsertBulk) UpdateAllowMessagesDispatch() *GroupUpsertBulk { }) } +// SetAllowLive sets the "allow_live" field. +func (u *GroupUpsertBulk) SetAllowLive(v bool) *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.SetAllowLive(v) + }) +} + +// UpdateAllowLive sets the "allow_live" field to the value that was provided on create. +func (u *GroupUpsertBulk) UpdateAllowLive() *GroupUpsertBulk { + return u.Update(func(s *GroupUpsert) { + s.UpdateAllowLive() + }) +} + // SetRequireOauthOnly sets the "require_oauth_only" field. func (u *GroupUpsertBulk) SetRequireOauthOnly(v bool) *GroupUpsertBulk { return u.Update(func(s *GroupUpsert) { diff --git a/backend/ent/group_update.go b/backend/ent/group_update.go index 52383d7668..595565df35 100644 --- a/backend/ent/group_update.go +++ b/backend/ent/group_update.go @@ -822,6 +822,20 @@ func (_u *GroupUpdate) SetNillableAllowMessagesDispatch(v *bool) *GroupUpdate { return _u } +// SetAllowLive sets the "allow_live" field. +func (_u *GroupUpdate) SetAllowLive(v bool) *GroupUpdate { + _u.mutation.SetAllowLive(v) + return _u +} + +// SetNillableAllowLive sets the "allow_live" field if the given value is not nil. +func (_u *GroupUpdate) SetNillableAllowLive(v *bool) *GroupUpdate { + if v != nil { + _u.SetAllowLive(*v) + } + return _u +} + // SetRequireOauthOnly sets the "require_oauth_only" field. func (_u *GroupUpdate) SetRequireOauthOnly(v bool) *GroupUpdate { _u.mutation.SetRequireOauthOnly(v) @@ -1495,6 +1509,9 @@ func (_u *GroupUpdate) sqlSave(ctx context.Context) (_node int, err error) { if value, ok := _u.mutation.AllowMessagesDispatch(); ok { _spec.SetField(group.FieldAllowMessagesDispatch, field.TypeBool, value) } + if value, ok := _u.mutation.AllowLive(); ok { + _spec.SetField(group.FieldAllowLive, field.TypeBool, value) + } if value, ok := _u.mutation.RequireOauthOnly(); ok { _spec.SetField(group.FieldRequireOauthOnly, field.TypeBool, value) } @@ -2627,6 +2644,20 @@ func (_u *GroupUpdateOne) SetNillableAllowMessagesDispatch(v *bool) *GroupUpdate return _u } +// SetAllowLive sets the "allow_live" field. +func (_u *GroupUpdateOne) SetAllowLive(v bool) *GroupUpdateOne { + _u.mutation.SetAllowLive(v) + return _u +} + +// SetNillableAllowLive sets the "allow_live" field if the given value is not nil. +func (_u *GroupUpdateOne) SetNillableAllowLive(v *bool) *GroupUpdateOne { + if v != nil { + _u.SetAllowLive(*v) + } + return _u +} + // SetRequireOauthOnly sets the "require_oauth_only" field. func (_u *GroupUpdateOne) SetRequireOauthOnly(v bool) *GroupUpdateOne { _u.mutation.SetRequireOauthOnly(v) @@ -3330,6 +3361,9 @@ func (_u *GroupUpdateOne) sqlSave(ctx context.Context) (_node *Group, err error) if value, ok := _u.mutation.AllowMessagesDispatch(); ok { _spec.SetField(group.FieldAllowMessagesDispatch, field.TypeBool, value) } + if value, ok := _u.mutation.AllowLive(); ok { + _spec.SetField(group.FieldAllowLive, field.TypeBool, value) + } if value, ok := _u.mutation.RequireOauthOnly(); ok { _spec.SetField(group.FieldRequireOauthOnly, field.TypeBool, value) } diff --git a/backend/ent/migrate/schema.go b/backend/ent/migrate/schema.go index d8600a87c6..5d78a74338 100644 --- a/backend/ent/migrate/schema.go +++ b/backend/ent/migrate/schema.go @@ -938,6 +938,7 @@ var ( {Name: "supported_model_scopes", Type: field.TypeJSON, SchemaType: map[string]string{"postgres": "jsonb"}}, {Name: "sort_order", Type: field.TypeInt, Default: 0}, {Name: "allow_messages_dispatch", Type: field.TypeBool, Default: false}, + {Name: "allow_live", Type: field.TypeBool, Default: false}, {Name: "require_oauth_only", Type: field.TypeBool, Default: false}, {Name: "require_privacy_set", Type: field.TypeBool, Default: false}, {Name: "default_mapped_model", Type: field.TypeString, Size: 100, Default: ""}, diff --git a/backend/ent/mutation.go b/backend/ent/mutation.go index d0633e5425..cf9258ffa9 100644 --- a/backend/ent/mutation.go +++ b/backend/ent/mutation.go @@ -21911,6 +21911,7 @@ type GroupMutation struct { sort_order *int addsort_order *int allow_messages_dispatch *bool + allow_live *bool require_oauth_only *bool require_privacy_set *bool default_mapped_model *string @@ -24226,6 +24227,42 @@ func (m *GroupMutation) ResetAllowMessagesDispatch() { m.allow_messages_dispatch = nil } +// SetAllowLive sets the "allow_live" field. +func (m *GroupMutation) SetAllowLive(b bool) { + m.allow_live = &b +} + +// AllowLive returns the value of the "allow_live" field in the mutation. +func (m *GroupMutation) AllowLive() (r bool, exists bool) { + v := m.allow_live + if v == nil { + return + } + return *v, true +} + +// OldAllowLive returns the old "allow_live" field's value of the Group entity. +// If the Group object wasn't provided to the builder, the object is fetched from the database. +// An error is returned if the mutation operation is not UpdateOne, or the database query fails. +func (m *GroupMutation) OldAllowLive(ctx context.Context) (v bool, err error) { + if !m.op.Is(OpUpdateOne) { + return v, errors.New("OldAllowLive is only allowed on UpdateOne operations") + } + if m.id == nil || m.oldValue == nil { + return v, errors.New("OldAllowLive requires an ID field in the mutation") + } + oldValue, err := m.oldValue(ctx) + if err != nil { + return v, fmt.Errorf("querying old value for OldAllowLive: %w", err) + } + return oldValue.AllowLive, nil +} + +// ResetAllowLive resets all changes to the "allow_live" field. +func (m *GroupMutation) ResetAllowLive() { + m.allow_live = nil +} + // SetRequireOauthOnly sets the "require_oauth_only" field. func (m *GroupMutation) SetRequireOauthOnly(b bool) { m.require_oauth_only = &b @@ -24907,7 +24944,7 @@ func (m *GroupMutation) Type() string { // order to get all numeric fields that were incremented/decremented, call // AddedFields(). func (m *GroupMutation) Fields() []string { - fields := make([]string, 0, 51) + fields := make([]string, 0, 52) if m.created_at != nil { fields = append(fields, group.FieldCreatedAt) } @@ -25037,6 +25074,9 @@ func (m *GroupMutation) Fields() []string { if m.allow_messages_dispatch != nil { fields = append(fields, group.FieldAllowMessagesDispatch) } + if m.allow_live != nil { + fields = append(fields, group.FieldAllowLive) + } if m.require_oauth_only != nil { fields = append(fields, group.FieldRequireOauthOnly) } @@ -25155,6 +25195,8 @@ func (m *GroupMutation) Field(name string) (ent.Value, bool) { return m.SortOrder() case group.FieldAllowMessagesDispatch: return m.AllowMessagesDispatch() + case group.FieldAllowLive: + return m.AllowLive() case group.FieldRequireOauthOnly: return m.RequireOauthOnly() case group.FieldRequirePrivacySet: @@ -25266,6 +25308,8 @@ func (m *GroupMutation) OldField(ctx context.Context, name string) (ent.Value, e return m.OldSortOrder(ctx) case group.FieldAllowMessagesDispatch: return m.OldAllowMessagesDispatch(ctx) + case group.FieldAllowLive: + return m.OldAllowLive(ctx) case group.FieldRequireOauthOnly: return m.OldRequireOauthOnly(ctx) case group.FieldRequirePrivacySet: @@ -25592,6 +25636,13 @@ func (m *GroupMutation) SetField(name string, value ent.Value) error { } m.SetAllowMessagesDispatch(v) return nil + case group.FieldAllowLive: + v, ok := value.(bool) + if !ok { + return fmt.Errorf("unexpected type %T for field %s", value, name) + } + m.SetAllowLive(v) + return nil case group.FieldRequireOauthOnly: v, ok := value.(bool) if !ok { @@ -26180,6 +26231,9 @@ func (m *GroupMutation) ResetField(name string) error { case group.FieldAllowMessagesDispatch: m.ResetAllowMessagesDispatch() return nil + case group.FieldAllowLive: + m.ResetAllowLive() + return nil case group.FieldRequireOauthOnly: m.ResetRequireOauthOnly() return nil diff --git a/backend/ent/runtime/runtime.go b/backend/ent/runtime/runtime.go index 9e5a0119e3..58d1a5f35e 100644 --- a/backend/ent/runtime/runtime.go +++ b/backend/ent/runtime/runtime.go @@ -1141,40 +1141,44 @@ func init() { groupDescAllowMessagesDispatch := groupFields[39].Descriptor() // group.DefaultAllowMessagesDispatch holds the default value on creation for the allow_messages_dispatch field. group.DefaultAllowMessagesDispatch = groupDescAllowMessagesDispatch.Default.(bool) + // groupDescAllowLive is the schema descriptor for allow_live field. + groupDescAllowLive := groupFields[40].Descriptor() + // group.DefaultAllowLive holds the default value on creation for the allow_live field. + group.DefaultAllowLive = groupDescAllowLive.Default.(bool) // groupDescRequireOauthOnly is the schema descriptor for require_oauth_only field. - groupDescRequireOauthOnly := groupFields[40].Descriptor() + groupDescRequireOauthOnly := groupFields[41].Descriptor() // group.DefaultRequireOauthOnly holds the default value on creation for the require_oauth_only field. group.DefaultRequireOauthOnly = groupDescRequireOauthOnly.Default.(bool) // groupDescRequirePrivacySet is the schema descriptor for require_privacy_set field. - groupDescRequirePrivacySet := groupFields[41].Descriptor() + groupDescRequirePrivacySet := groupFields[42].Descriptor() // group.DefaultRequirePrivacySet holds the default value on creation for the require_privacy_set field. group.DefaultRequirePrivacySet = groupDescRequirePrivacySet.Default.(bool) // groupDescDefaultMappedModel is the schema descriptor for default_mapped_model field. - groupDescDefaultMappedModel := groupFields[42].Descriptor() + groupDescDefaultMappedModel := groupFields[43].Descriptor() // group.DefaultDefaultMappedModel holds the default value on creation for the default_mapped_model field. group.DefaultDefaultMappedModel = groupDescDefaultMappedModel.Default.(string) // group.DefaultMappedModelValidator is a validator for the "default_mapped_model" field. It is called by the builders before save. group.DefaultMappedModelValidator = groupDescDefaultMappedModel.Validators[0].(func(string) error) // groupDescMessagesDispatchModelConfig is the schema descriptor for messages_dispatch_model_config field. - groupDescMessagesDispatchModelConfig := groupFields[43].Descriptor() + groupDescMessagesDispatchModelConfig := groupFields[44].Descriptor() // group.DefaultMessagesDispatchModelConfig holds the default value on creation for the messages_dispatch_model_config field. group.DefaultMessagesDispatchModelConfig = groupDescMessagesDispatchModelConfig.Default.(domain.OpenAIMessagesDispatchModelConfig) // groupDescModelsListConfig is the schema descriptor for models_list_config field. - groupDescModelsListConfig := groupFields[44].Descriptor() + groupDescModelsListConfig := groupFields[45].Descriptor() // group.DefaultModelsListConfig holds the default value on creation for the models_list_config field. group.DefaultModelsListConfig = groupDescModelsListConfig.Default.(domain.GroupModelsListConfig) // groupDescRpmLimit is the schema descriptor for rpm_limit field. - groupDescRpmLimit := groupFields[45].Descriptor() + groupDescRpmLimit := groupFields[46].Descriptor() // group.DefaultRpmLimit holds the default value on creation for the rpm_limit field. group.DefaultRpmLimit = groupDescRpmLimit.Default.(int) // groupDescMaxReasoningEffort is the schema descriptor for max_reasoning_effort field. - groupDescMaxReasoningEffort := groupFields[46].Descriptor() + groupDescMaxReasoningEffort := groupFields[47].Descriptor() // group.DefaultMaxReasoningEffort holds the default value on creation for the max_reasoning_effort field. group.DefaultMaxReasoningEffort = groupDescMaxReasoningEffort.Default.(string) // group.MaxReasoningEffortValidator is a validator for the "max_reasoning_effort" field. It is called by the builders before save. group.MaxReasoningEffortValidator = groupDescMaxReasoningEffort.Validators[0].(func(string) error) // groupDescReasoningEffortMappings is the schema descriptor for reasoning_effort_mappings field. - groupDescReasoningEffortMappings := groupFields[47].Descriptor() + groupDescReasoningEffortMappings := groupFields[48].Descriptor() // group.DefaultReasoningEffortMappings holds the default value on creation for the reasoning_effort_mappings field. group.DefaultReasoningEffortMappings = groupDescReasoningEffortMappings.Default.([]domain.ReasoningEffortMapping) idempotencyrecordMixin := schema.IdempotencyRecord{}.Mixin() diff --git a/backend/ent/schema/group.go b/backend/ent/schema/group.go index bce09c1f8a..4b6d056ef6 100644 --- a/backend/ent/schema/group.go +++ b/backend/ent/schema/group.go @@ -198,6 +198,9 @@ func (Group) Fields() []ent.Field { field.Bool("allow_messages_dispatch"). Default(false). Comment("是否允许 /v1/messages 调度到此 OpenAI 分组"), + field.Bool("allow_live"). + Default(false). + Comment("是否允许此 OpenAI 分组访问 Live 接口"), field.Bool("require_oauth_only"). Default(false). Comment("仅允许非 apikey 类型账号关联到此分组"), diff --git a/backend/internal/config/config.go b/backend/internal/config/config.go index 95bc3a0580..34f02f9898 100644 --- a/backend/internal/config/config.go +++ b/backend/internal/config/config.go @@ -913,6 +913,8 @@ type GatewayConfig struct { OpenAICompactModel string `mapstructure:"openai_compact_model"` // OpenAIWS: OpenAI Responses WebSocket 配置(默认开启,可按需回滚到 HTTP) OpenAIWS GatewayOpenAIWSConfig `mapstructure:"openai_ws"` + // Live: ChatGPT Frameless Live 会话配置。 + Live GatewayLiveConfig `mapstructure:"live"` // OpenAIScheduler: OpenAI 高级调度器粘性逃逸配置 OpenAIScheduler GatewayOpenAISchedulerConfig `mapstructure:"openai_scheduler"` // OpenAIHTTP2: OpenAI HTTP 上游协议策略(默认启用 HTTP/2,可按代理能力回退 HTTP/1.1) @@ -999,6 +1001,11 @@ type GatewayConfig struct { UserMessageQueue UserMessageQueueConfig `mapstructure:"user_message_queue"` } +type GatewayLiveConfig struct { + // MaxSessionDurationSeconds 是 Live 会话的硬上限。 + MaxSessionDurationSeconds int `mapstructure:"max_session_duration_seconds"` +} + // GatewayOpenAIHTTP2Config OpenAI HTTP 上游协议配置。 // 默认启用 HTTP/2;在部分代理不兼容时按策略回退 HTTP/1.1。 type GatewayOpenAIHTTP2Config struct { @@ -3039,6 +3046,9 @@ func (c *Config) Validate() error { (c.Gateway.OpenAIHighEffortFirstOutputTimeoutSeconds > 0 && c.Gateway.OpenAIHighEffortFirstOutputTimeoutSeconds < 30) { return fmt.Errorf("gateway.openai_high_effort_first_output_timeout_seconds must be 0 or between 30-1800 seconds") } + if c.Gateway.Live.MaxSessionDurationSeconds <= 0 { + c.Gateway.Live.MaxSessionDurationSeconds = 3600 + } if strings.TrimSpace(c.Gateway.ConnectionPoolIsolation) != "" { switch c.Gateway.ConnectionPoolIsolation { case ConnectionPoolIsolationProxy, ConnectionPoolIsolationAccount, ConnectionPoolIsolationAccountProxy: diff --git a/backend/internal/handler/admin/group_handler.go b/backend/internal/handler/admin/group_handler.go index 070582663f..b913187638 100644 --- a/backend/internal/handler/admin/group_handler.go +++ b/backend/internal/handler/admin/group_handler.go @@ -125,6 +125,7 @@ type CreateGroupRequest struct { SupportedModelScopes []string `json:"supported_model_scopes"` // OpenAI Messages 调度配置(仅 openai 平台使用) AllowMessagesDispatch bool `json:"allow_messages_dispatch"` + AllowLive bool `json:"allow_live"` RequireOAuthOnly bool `json:"require_oauth_only"` RequirePrivacySet bool `json:"require_privacy_set"` DefaultMappedModel string `json:"default_mapped_model"` @@ -183,6 +184,7 @@ type UpdateGroupRequest struct { SupportedModelScopes *[]string `json:"supported_model_scopes"` // OpenAI Messages 调度配置(仅 openai 平台使用) AllowMessagesDispatch *bool `json:"allow_messages_dispatch"` + AllowLive *bool `json:"allow_live"` RequireOAuthOnly *bool `json:"require_oauth_only"` RequirePrivacySet *bool `json:"require_privacy_set"` DefaultMappedModel *string `json:"default_mapped_model"` @@ -499,6 +501,7 @@ func (h *GroupHandler) Create(c *gin.Context) { MCPXMLInject: req.MCPXMLInject, SupportedModelScopes: req.SupportedModelScopes, AllowMessagesDispatch: req.AllowMessagesDispatch, + AllowLive: req.AllowLive, RequireOAuthOnly: req.RequireOAuthOnly, RequirePrivacySet: req.RequirePrivacySet, DefaultMappedModel: req.DefaultMappedModel, @@ -617,6 +620,7 @@ func (h *GroupHandler) Update(c *gin.Context) { MCPXMLInject: req.MCPXMLInject, SupportedModelScopes: req.SupportedModelScopes, AllowMessagesDispatch: req.AllowMessagesDispatch, + AllowLive: req.AllowLive, RequireOAuthOnly: req.RequireOAuthOnly, RequirePrivacySet: req.RequirePrivacySet, DefaultMappedModel: req.DefaultMappedModel, diff --git a/backend/internal/handler/dto/mappers.go b/backend/internal/handler/dto/mappers.go index 1ae9a5e790..567d12c13a 100644 --- a/backend/internal/handler/dto/mappers.go +++ b/backend/internal/handler/dto/mappers.go @@ -204,6 +204,7 @@ func groupFromServiceBase(g *service.Group) Group { FallbackGroupID: g.FallbackGroupID, FallbackGroupIDOnInvalidRequest: g.FallbackGroupIDOnInvalidRequest, AllowMessagesDispatch: g.AllowMessagesDispatch, + AllowLive: g.AllowLive, RequireOAuthOnly: g.RequireOAuthOnly, RequirePrivacySet: g.RequirePrivacySet, RPMLimit: g.RPMLimit, diff --git a/backend/internal/handler/dto/types.go b/backend/internal/handler/dto/types.go index 38133cff3b..d52f81afe3 100644 --- a/backend/internal/handler/dto/types.go +++ b/backend/internal/handler/dto/types.go @@ -132,6 +132,8 @@ type Group struct { // OpenAI Messages 调度开关(用户侧需要此字段判断是否展示 Claude Code 教程) AllowMessagesDispatch bool `json:"allow_messages_dispatch"` + // OpenAI Live 接口开关 + AllowLive bool `json:"allow_live"` // 账号过滤控制(仅 OpenAI/Antigravity 平台有效) RequireOAuthOnly bool `json:"require_oauth_only"` diff --git a/backend/internal/handler/openai_live.go b/backend/internal/handler/openai_live.go new file mode 100644 index 0000000000..466d84d685 --- /dev/null +++ b/backend/internal/handler/openai_live.go @@ -0,0 +1,231 @@ +package handler + +import ( + "encoding/json" + "errors" + "io" + "net/http" + "net/url" + "strconv" + "strings" + + "github.com/Wei-Shaw/sub2api/internal/pkg/ip" + middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware" + "github.com/Wei-Shaw/sub2api/internal/service" + coderws "github.com/coder/websocket" + "github.com/gin-gonic/gin" + "github.com/tidwall/gjson" + "go.uber.org/zap" +) + +func (h *OpenAIGatewayHandler) Live(c *gin.Context) { + apiKey, ok := middleware2.GetAPIKeyFromContext(c) + if !ok { + h.errorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key") + return + } + subject, ok := middleware2.GetAuthSubjectFromContext(c) + if !ok { + h.errorResponse(c, http.StatusInternalServerError, "api_error", "User context not found") + return + } + if apiKey.Group == nil || apiKey.Group.Platform != service.PlatformOpenAI { + h.errorResponse(c, http.StatusNotFound, "not_found_error", "Live is not supported for this platform") + return + } + if !liveEnabledForAPIKey(apiKey) { + h.errorResponse(c, http.StatusForbidden, "permission_error", "Live is not enabled for this group") + return + } + request, err := parseLiveCallRequest(c) + if err != nil { + h.errorResponse(c, http.StatusBadRequest, "invalid_request_error", err.Error()) + return + } + model := strings.TrimSpace(gjson.GetBytes(request.Session, "model").String()) + reqLog := requestLogger( + c, + "handler.openai_gateway.live", + zap.Int64("user_id", subject.UserID), + zap.Int64("api_key_id", apiKey.ID), + zap.Any("group_id", apiKey.GroupID), + ) + if decision := h.checkSecurityAudit( + c, + reqLog, + apiKey, + subject, + service.ContentModerationProtocolOpenAIResponses, + model, + request.Session, + ); decision != nil && !decision.AllowNextStage { + h.openAISecurityAuditError(c, decision) + return + } + + subscription, _ := middleware2.GetSubscriptionFromContext(c) + if h.billingCacheService == nil { + h.errorResponse(c, http.StatusServiceUnavailable, "api_error", "Billing service unavailable") + return + } + if err := h.billingCacheService.CheckBillingEligibility( + c.Request.Context(), + apiKey.User, + apiKey, + apiKey.Group, + subscription, + service.QuotaPlatform(c.Request.Context(), apiKey), + ); err != nil { + status, code, message, retryAfter := billingErrorDetails(err) + if retryAfter > 0 { + c.Header("Retry-After", strconv.Itoa(retryAfter)) + } + h.errorResponse(c, status, code, message) + return + } + + userRelease, acquired, err := h.concurrencyHelper.TryAcquireUserSlot( + c.Request.Context(), + subject.UserID, + subject.Concurrency, + ) + if err != nil { + h.errorResponse(c, http.StatusServiceUnavailable, "api_error", "Live concurrency unavailable") + return + } + if !acquired { + h.errorResponse(c, http.StatusTooManyRequests, "rate_limit_error", "Live concurrency limit reached") + return + } + defer userRelease() + + identity := liveCallIdentity(c, apiKey, subject.UserID, subscription) + created, err := h.gatewayService.CreateLiveCall(c.Request.Context(), request, identity, subject.Concurrency) + if err != nil { + h.writeLiveCreateError(c, err) + return + } + c.Header("Location", liveSidebandLocation(c.FullPath(), created.CallID)) + c.Data(http.StatusOK, "application/sdp", created.SDP) +} + +func parseLiveCallRequest(c *gin.Context) (*service.LiveCallRequest, error) { + contentType := strings.ToLower(c.GetHeader("Content-Type")) + if strings.HasPrefix(contentType, "multipart/form-data") { + sdp := c.PostForm("sdp") + session := json.RawMessage(c.PostForm("session")) + request := &service.LiveCallRequest{SDP: sdp, Session: session} + if err := service.ValidateLiveCallRequest(request); err != nil { + return nil, err + } + return request, nil + } + var request service.LiveCallRequest + decoder := json.NewDecoder(c.Request.Body) + if err := decoder.Decode(&request); err != nil { + return nil, errors.New("request body must be valid JSON") + } + if err := decoder.Decode(&struct{}{}); !errors.Is(err, io.EOF) { + return nil, errors.New("request body must contain one JSON object") + } + if err := service.ValidateLiveCallRequest(&request); err != nil { + return nil, err + } + return &request, nil +} + +func liveSidebandLocation(fullPath, callID string) string { + prefix := "/v1/live/" + if strings.HasPrefix(fullPath, "/backend-api/codex/") { + prefix = "/backend-api/codex/" + } + return prefix + url.PathEscape(callID) +} + +func liveCallIdentity( + c *gin.Context, + apiKey *service.APIKey, + userID int64, + subscription *service.UserSubscription, +) service.LiveCallIdentity { + var subscriptionID *int64 + if subscription != nil { + value := subscription.ID + subscriptionID = &value + } + return service.LiveCallIdentity{ + APIKeyID: apiKey.ID, + UserID: userID, + GroupID: apiKey.GroupID, + SubscriptionID: subscriptionID, + UserAgent: c.GetHeader("User-Agent"), + IPAddress: ip.GetClientIP(c), + InboundEndpoint: GetInboundEndpoint(c), + } +} + +func (h *OpenAIGatewayHandler) writeLiveCreateError(c *gin.Context, err error) { + switch { + case errors.Is(err, service.ErrLiveConcurrencyFull): + h.errorResponse(c, http.StatusTooManyRequests, "rate_limit_error", "Live concurrency limit reached") + case errors.Is(err, service.ErrLiveUnavailable): + h.errorResponse(c, http.StatusServiceUnavailable, "api_error", "Live is unavailable") + default: + var upstreamErr *service.UpstreamFailoverError + if errors.As(err, &upstreamErr) && upstreamErr.StatusCode >= 400 && upstreamErr.StatusCode < 500 { + h.errorResponse(c, upstreamErr.StatusCode, "invalid_request_error", "Live upstream rejected the request") + return + } + h.errorResponse(c, http.StatusBadGateway, "api_error", "Live upstream request failed") + } +} + +func (h *OpenAIGatewayHandler) LiveSideband(c *gin.Context) { + apiKey, ok := middleware2.GetAPIKeyFromContext(c) + if !ok { + h.errorResponse(c, http.StatusUnauthorized, "authentication_error", "Invalid API key") + return + } + subject, ok := middleware2.GetAuthSubjectFromContext(c) + if !ok { + h.errorResponse(c, http.StatusInternalServerError, "api_error", "User context not found") + return + } + if !liveEnabledForAPIKey(apiKey) { + h.errorResponse(c, http.StatusForbidden, "permission_error", "Live is not enabled for this group") + return + } + identity := service.LiveCallIdentity{ + APIKeyID: apiKey.ID, + UserID: subject.UserID, + GroupID: apiKey.GroupID, + } + record, err := h.gatewayService.GetLiveCallForIdentity(c.Request.Context(), c.Param("call_id"), identity) + if err != nil { + if errors.Is(err, service.ErrLiveIdentityMismatch) { + h.errorResponse(c, http.StatusForbidden, "permission_error", "Live call belongs to another identity") + return + } + h.errorResponse(c, http.StatusNotFound, "not_found_error", "Live call not found") + return + } + downstream, err := coderws.Accept(c.Writer, c.Request, &coderws.AcceptOptions{ + InsecureSkipVerify: true, + }) + if err != nil { + return + } + defer downstream.CloseNow() + if err := h.gatewayService.ProxyLiveSideband(c.Request.Context(), record, downstream); err != nil { + _ = downstream.Close(coderws.StatusInternalError, "live sideband closed") + return + } + _ = downstream.Close(coderws.StatusNormalClosure, "") +} + +func liveEnabledForAPIKey(apiKey *service.APIKey) bool { + return apiKey != nil && + apiKey.Group != nil && + apiKey.Group.Platform == service.PlatformOpenAI && + apiKey.Group.AllowLive +} diff --git a/backend/internal/handler/openai_live_test.go b/backend/internal/handler/openai_live_test.go new file mode 100644 index 0000000000..dde35f9f98 --- /dev/null +++ b/backend/internal/handler/openai_live_test.go @@ -0,0 +1,104 @@ +package handler + +import ( + "bytes" + "encoding/json" + "mime/multipart" + "net/http/httptest" + "testing" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestParseLiveCallRequestMultipartPreservesSession(t *testing.T) { + gin.SetMode(gin.TestMode) + session := `{"model":"gpt-live-test","delegation":{"type":"client"},"instructions":"你好"}` + var body bytes.Buffer + writer := multipart.NewWriter(&body) + require.NoError(t, writer.WriteField("sdp", "v=0\r\n")) + require.NoError(t, writer.WriteField("session", session)) + require.NoError(t, writer.Close()) + + request := httptest.NewRequest("POST", "/v1/live", &body) + request.Header.Set("Content-Type", writer.FormDataContentType()) + context, _ := gin.CreateTestContext(httptest.NewRecorder()) + context.Request = request + + parsed, err := parseLiveCallRequest(context) + require.NoError(t, err) + require.Equal(t, "v=0\r\n", parsed.SDP) + require.JSONEq(t, session, string(parsed.Session)) + require.Equal(t, "client", jsonPathString(t, parsed.Session, "delegation", "type")) +} + +func TestParseLiveCallRequestJSONPreservesSessionWithoutDelegation(t *testing.T) { + gin.SetMode(gin.TestMode) + body := `{"sdp":"v=0\\r\\n","session":{"model":"gpt-live-test","instructions":"standalone"}}` + request := httptest.NewRequest("POST", "/backend-api/codex/realtime/calls", bytes.NewBufferString(body)) + request.Header.Set("Content-Type", "application/json") + context, _ := gin.CreateTestContext(httptest.NewRecorder()) + context.Request = request + + parsed, err := parseLiveCallRequest(context) + require.NoError(t, err) + require.NotContains(t, string(parsed.Session), "delegation") + require.Equal(t, "standalone", jsonPathString(t, parsed.Session, "instructions")) +} + +func TestParseLiveCallRequestRejectsInvalidJSONShape(t *testing.T) { + gin.SetMode(gin.TestMode) + testCases := []string{ + `{"session":{"type":"quicksilver"}}`, + `{"sdp":"v=0\\r\\n","session":[]}`, + `{"sdp":"v=0\\r\\n","session":null}`, + `{"sdp":"v=0\\r\\n","session":{"type":"quicksilver"}} {}`, + } + for _, body := range testCases { + request := httptest.NewRequest("POST", "/backend-api/codex/realtime/calls", bytes.NewBufferString(body)) + request.Header.Set("Content-Type", "application/json") + context, _ := gin.CreateTestContext(httptest.NewRecorder()) + context.Request = request + _, err := parseLiveCallRequest(context) + require.Error(t, err) + } +} + +func TestLiveSidebandLocationMatchesCreateRoute(t *testing.T) { + require.Equal(t, "/v1/live/call_123", liveSidebandLocation("/v1/live", "call_123")) + require.Equal( + t, + "/backend-api/codex/call_123", + liveSidebandLocation("/backend-api/codex/realtime/calls", "call_123"), + ) +} + +func TestLiveEnabledForAPIKey(t *testing.T) { + require.False(t, liveEnabledForAPIKey(nil)) + require.False(t, liveEnabledForAPIKey(&service.APIKey{})) + require.False(t, liveEnabledForAPIKey(&service.APIKey{ + Group: &service.Group{Platform: service.PlatformOpenAI}, + })) + require.False(t, liveEnabledForAPIKey(&service.APIKey{ + Group: &service.Group{Platform: service.PlatformAnthropic, AllowLive: true}, + })) + require.True(t, liveEnabledForAPIKey(&service.APIKey{ + Group: &service.Group{Platform: service.PlatformOpenAI, AllowLive: true}, + })) +} + +func jsonPathString(t *testing.T, raw json.RawMessage, keys ...string) string { + t.Helper() + var value any + require.NoError(t, json.Unmarshal(raw, &value)) + current := value + for _, key := range keys { + object, ok := current.(map[string]any) + require.True(t, ok) + current = object[key] + } + result, ok := current.(string) + require.True(t, ok) + return result +} diff --git a/backend/internal/repository/api_key_repo.go b/backend/internal/repository/api_key_repo.go index 8addad4342..2b66b91d76 100644 --- a/backend/internal/repository/api_key_repo.go +++ b/backend/internal/repository/api_key_repo.go @@ -199,6 +199,7 @@ func (r *apiKeyRepository) GetByKeyForAuth(ctx context.Context, key string) (*se group.FieldMcpXMLInject, group.FieldSupportedModelScopes, group.FieldAllowMessagesDispatch, + group.FieldAllowLive, group.FieldDefaultMappedModel, group.FieldMessagesDispatchModelConfig, group.FieldModelsListConfig, @@ -948,6 +949,7 @@ func groupEntityToService(g *dbent.Group) *service.Group { SupportedModelScopes: g.SupportedModelScopes, SortOrder: g.SortOrder, AllowMessagesDispatch: g.AllowMessagesDispatch, + AllowLive: g.AllowLive, RequireOAuthOnly: g.RequireOauthOnly, RequirePrivacySet: g.RequirePrivacySet, DefaultMappedModel: g.DefaultMappedModel, diff --git a/backend/internal/repository/concurrency_cache.go b/backend/internal/repository/concurrency_cache.go index 5341d411e9..37e839f287 100644 --- a/backend/internal/repository/concurrency_cache.go +++ b/backend/internal/repository/concurrency_cache.go @@ -29,11 +29,15 @@ const ( // 格式: concurrency:user:{userID} userSlotKeyPrefix = "concurrency:user:" // 格式: concurrency:api_key:{apiKeyID} - apiKeySlotKeyPrefix = "concurrency:api_key:" + apiKeySlotKeyPrefix = "concurrency:api_key:" + liveAccountSlotKeyPrefix = "concurrency:live:account:" + liveUserSlotKeyPrefix = "concurrency:live:user:" + liveAPIKeySlotKeyPrefix = "concurrency:live:api_key:" // API-key-scoped client WebSocket ingress leases use a shorter TTL than // ordinary request slots, because idle ingress sessions do not hold a turn slot. openAIWSIngressLeaseKeyPrefix = "concurrency:openai_ws_ingress:api_key:" openAIWSIngressLeaseTTLSeconds = 60 + liveLeaseTTLSeconds = 60 // 等待队列计数器格式: concurrency:wait:{userID} waitQueueKeyPrefix = "concurrency:wait:" // 账号级等待队列计数器格式: wait:account:{accountID} @@ -59,7 +63,7 @@ const ( var ( // acquireScript 使用有序集合计数并在未达上限时添加槽位 // 使用 Redis TIME 命令获取服务器时间,避免多实例时钟不同步问题 - // KEYS[1] = 有序集合键 (concurrency:account:{id} / concurrency:user:{id}) + // KEYS[1] = 普通槽位键,KEYS[2] = 对应 Live 槽位键 // ARGV[1] = maxConcurrency // ARGV[2] = TTL(秒) // ARGV[3] = requestID @@ -69,6 +73,7 @@ var ( -- replicates correctly. No-op on Redis 5.0+ (effects replication is default). redis.replicate_commands() local key = KEYS[1] + local liveKey = KEYS[2] local maxConcurrency = tonumber(ARGV[1]) local ttl = tonumber(ARGV[2]) local requestID = ARGV[3] @@ -80,6 +85,7 @@ var ( -- 清理过期槽位 redis.call('ZREMRANGEBYSCORE', key, '-inf', expireBefore) + redis.call('ZREMRANGEBYSCORE', liveKey, '-inf', now - 60) -- 检查是否已存在(支持重试场景刷新时间戳) local exists = redis.call('ZSCORE', key, requestID) @@ -90,7 +96,7 @@ var ( end -- 检查是否达到并发上限 - local count = redis.call('ZCARD', key) + local count = redis.call('ZCARD', key) + redis.call('ZCARD', liveKey) if count < maxConcurrency then redis.call('ZADD', key, now, requestID) redis.call('EXPIRE', key, ttl) @@ -102,13 +108,14 @@ var ( // getCountScript 统计有序集合中的槽位数量并清理过期条目 // 使用 Redis TIME 命令获取服务器时间 - // KEYS[1] = 有序集合键 + // KEYS[1] = 普通槽位键,KEYS[2] = 对应 Live 槽位键 // ARGV[1] = TTL(秒) getCountScript = redis.NewScript(` -- Redis 3.2-4.x compat: opt into effects replication so redis.call('TIME') -- replicates correctly. No-op on Redis 5.0+ (effects replication is default). redis.replicate_commands() local key = KEYS[1] + local liveKey = KEYS[2] local ttl = tonumber(ARGV[1]) -- 使用 Redis 服务器时间 @@ -117,7 +124,60 @@ var ( local expireBefore = now - ttl redis.call('ZREMRANGEBYSCORE', key, '-inf', expireBefore) - return redis.call('ZCARD', key) + redis.call('ZREMRANGEBYSCORE', liveKey, '-inf', now - 60) + return redis.call('ZCARD', key) + redis.call('ZCARD', liveKey) + `) + + acquireLiveLeaseScript = redis.NewScript(` + redis.replicate_commands() + local accountRegular = KEYS[1] + local accountLive = KEYS[2] + local userRegular = KEYS[3] + local userLive = KEYS[4] + local apiLive = KEYS[5] + local accountMax = tonumber(ARGV[1]) + local userMax = tonumber(ARGV[2]) + local ttl = tonumber(ARGV[3]) + local leaseID = ARGV[4] + local replacing = tonumber(ARGV[5]) + local now = tonumber(redis.call('TIME')[1]) + local liveExpireBefore = now - ttl + redis.call('ZREMRANGEBYSCORE', accountLive, '-inf', liveExpireBefore) + redis.call('ZREMRANGEBYSCORE', userLive, '-inf', liveExpireBefore) + redis.call('ZREMRANGEBYSCORE', apiLive, '-inf', liveExpireBefore) + if redis.call('ZSCORE', accountLive, leaseID) ~= false then + return 1 + end + local accountCount = redis.call('ZCARD', accountRegular) + redis.call('ZCARD', accountLive) + local userCount = redis.call('ZCARD', userRegular) + redis.call('ZCARD', userLive) + local allowance = 0 + if replacing == 1 then allowance = 1 end + if accountMax > 0 and accountCount >= accountMax + allowance then return 0 end + if userMax > 0 and userCount >= userMax + allowance then return 0 end + redis.call('ZADD', accountLive, now, leaseID) + redis.call('ZADD', userLive, now, leaseID) + redis.call('ZADD', apiLive, now, leaseID) + redis.call('EXPIRE', accountLive, ttl) + redis.call('EXPIRE', userLive, ttl) + redis.call('EXPIRE', apiLive, ttl) + return 1 + `) + + refreshLiveLeaseScript = redis.NewScript(` + redis.replicate_commands() + local ttl = tonumber(ARGV[1]) + local leaseID = ARGV[2] + local now = tonumber(redis.call('TIME')[1]) + local expireBefore = now - ttl + for _, key in ipairs(KEYS) do + redis.call('ZREMRANGEBYSCORE', key, '-inf', expireBefore) + if redis.call('ZSCORE', key, leaseID) == false then return 0 end + end + for _, key in ipairs(KEYS) do + redis.call('ZADD', key, now, leaseID) + redis.call('EXPIRE', key, ttl) + end + return 1 `) // trackSlotScript 记录 stats-only 槽位,不做并发上限判断。 @@ -330,6 +390,18 @@ func apiKeySlotKey(apiKeyID int64) string { return fmt.Sprintf("%s%d", apiKeySlotKeyPrefix, apiKeyID) } +func liveAccountSlotKey(accountID int64) string { + return fmt.Sprintf("%s%d", liveAccountSlotKeyPrefix, accountID) +} + +func liveUserSlotKey(userID int64) string { + return fmt.Sprintf("%s%d", liveUserSlotKeyPrefix, userID) +} + +func liveAPIKeySlotKey(apiKeyID int64) string { + return fmt.Sprintf("%s%d", liveAPIKeySlotKeyPrefix, apiKeyID) +} + func openAIWSIngressLeaseKey(apiKeyID int64) string { return fmt.Sprintf("%s%d", openAIWSIngressLeaseKeyPrefix, apiKeyID) } @@ -559,7 +631,7 @@ func runScriptInt64Pair(ctx context.Context, rdb *redis.Client, script *redis.Sc func (c *concurrencyCache) AcquireAccountSlot(ctx context.Context, accountID int64, maxConcurrency int, requestID string) (bool, error) { key := accountSlotKey(accountID) // 时间戳在 Lua 脚本内使用 Redis TIME 命令获取,确保多实例时钟一致 - result, now, err := runScriptInt64Pair(ctx, c.rdb, acquireScript, []string{key}, maxConcurrency, c.slotTTLSeconds, requestID) + result, now, err := runScriptInt64Pair(ctx, c.rdb, acquireScript, []string{key, liveAccountSlotKey(accountID)}, maxConcurrency, c.slotTTLSeconds, requestID) if err != nil { return false, err } @@ -583,7 +655,7 @@ func (c *concurrencyCache) ReleaseAccountSlot(ctx context.Context, accountID int func (c *concurrencyCache) GetAccountConcurrency(ctx context.Context, accountID int64) (int, error) { key := accountSlotKey(accountID) // 时间戳在 Lua 脚本内使用 Redis TIME 命令获取 - result, err := getCountScript.Run(ctx, c.rdb, []string{key}, c.slotTTLSeconds).Int() + result, err := getCountScript.Run(ctx, c.rdb, []string{key, liveAccountSlotKey(accountID)}, c.slotTTLSeconds).Int() if err != nil { return 0, err } @@ -605,14 +677,18 @@ func (c *concurrencyCache) GetAccountConcurrencyBatch(ctx context.Context, accou type accountCmd struct { accountID int64 zcardCmd *redis.IntCmd + liveCmd *redis.IntCmd } cmds := make([]accountCmd, 0, len(accountIDs)) for _, accountID := range accountIDs { slotKey := accountSlotKeyPrefix + strconv.FormatInt(accountID, 10) + liveKey := liveAccountSlotKeyPrefix + strconv.FormatInt(accountID, 10) pipe.ZRemRangeByScore(ctx, slotKey, "-inf", strconv.FormatInt(cutoffTime, 10)) + pipe.ZRemRangeByScore(ctx, liveKey, "-inf", strconv.FormatInt(now.Unix()-liveLeaseTTLSeconds, 10)) cmds = append(cmds, accountCmd{ accountID: accountID, zcardCmd: pipe.ZCard(ctx, slotKey), + liveCmd: pipe.ZCard(ctx, liveKey), }) } @@ -622,7 +698,7 @@ func (c *concurrencyCache) GetAccountConcurrencyBatch(ctx context.Context, accou result := make(map[int64]int, len(accountIDs)) for _, cmd := range cmds { - result[cmd.accountID] = int(cmd.zcardCmd.Val()) + result[cmd.accountID] = int(cmd.zcardCmd.Val() + cmd.liveCmd.Val()) } return result, nil } @@ -632,7 +708,7 @@ func (c *concurrencyCache) GetAccountConcurrencyBatch(ctx context.Context, accou func (c *concurrencyCache) AcquireUserSlot(ctx context.Context, userID int64, maxConcurrency int, requestID string) (bool, error) { key := userSlotKey(userID) // 时间戳在 Lua 脚本内使用 Redis TIME 命令获取,确保多实例时钟一致 - result, now, err := runScriptInt64Pair(ctx, c.rdb, acquireScript, []string{key}, maxConcurrency, c.slotTTLSeconds, requestID) + result, now, err := runScriptInt64Pair(ctx, c.rdb, acquireScript, []string{key, liveUserSlotKey(userID)}, maxConcurrency, c.slotTTLSeconds, requestID) if err != nil { return false, err } @@ -656,7 +732,7 @@ func (c *concurrencyCache) ReleaseUserSlot(ctx context.Context, userID int64, re func (c *concurrencyCache) GetUserConcurrency(ctx context.Context, userID int64) (int, error) { key := userSlotKey(userID) // 时间戳在 Lua 脚本内使用 Redis TIME 命令获取 - result, err := getCountScript.Run(ctx, c.rdb, []string{key}, c.slotTTLSeconds).Int() + result, err := getCountScript.Run(ctx, c.rdb, []string{key, liveUserSlotKey(userID)}, c.slotTTLSeconds).Int() if err != nil { return 0, err } @@ -716,6 +792,57 @@ func (c *concurrencyCache) ReleaseOpenAIWSIngressLease(ctx context.Context, apiK return c.rdb.ZRem(ctx, openAIWSIngressLeaseKey(apiKeyID), leaseID).Err() } +func (c *concurrencyCache) AcquireLiveLease( + ctx context.Context, + accountID int64, + accountMax int, + userID int64, + userMax int, + apiKeyID int64, + leaseID string, + replacingRegularSlots bool, +) (bool, error) { + if c == nil || c.rdb == nil || accountID <= 0 || userID <= 0 || apiKeyID <= 0 || leaseID == "" { + return false, nil + } + replacing := 0 + if replacingRegularSlots { + replacing = 1 + } + result, err := acquireLiveLeaseScript.Run(ctx, c.rdb, []string{ + accountSlotKey(accountID), + liveAccountSlotKey(accountID), + userSlotKey(userID), + liveUserSlotKey(userID), + liveAPIKeySlotKey(apiKeyID), + }, accountMax, userMax, liveLeaseTTLSeconds, leaseID, replacing).Int() + return result == 1, err +} + +func (c *concurrencyCache) RefreshLiveLease(ctx context.Context, accountID, userID, apiKeyID int64, leaseID string) (bool, error) { + if c == nil || c.rdb == nil || leaseID == "" { + return false, nil + } + result, err := refreshLiveLeaseScript.Run(ctx, c.rdb, []string{ + liveAccountSlotKey(accountID), + liveUserSlotKey(userID), + liveAPIKeySlotKey(apiKeyID), + }, liveLeaseTTLSeconds, leaseID).Int() + return result == 1, err +} + +func (c *concurrencyCache) ReleaseLiveLease(ctx context.Context, accountID, userID, apiKeyID int64, leaseID string) error { + if c == nil || c.rdb == nil || leaseID == "" { + return nil + } + pipe := c.rdb.TxPipeline() + pipe.ZRem(ctx, liveAccountSlotKey(accountID), leaseID) + pipe.ZRem(ctx, liveUserSlotKey(userID), leaseID) + pipe.ZRem(ctx, liveAPIKeySlotKey(apiKeyID), leaseID) + _, err := pipe.Exec(ctx) + return err +} + func (c *concurrencyCache) GetAPIKeyConcurrencyBatch(ctx context.Context, apiKeyIDs []int64) (map[int64]int, error) { if len(apiKeyIDs) == 0 { return map[int64]int{}, nil @@ -731,14 +858,18 @@ func (c *concurrencyCache) GetAPIKeyConcurrencyBatch(ctx context.Context, apiKey type apiKeyCmd struct { apiKeyID int64 zcardCmd *redis.IntCmd + liveCmd *redis.IntCmd } cmds := make([]apiKeyCmd, 0, len(apiKeyIDs)) for _, apiKeyID := range apiKeyIDs { slotKey := apiKeySlotKeyPrefix + strconv.FormatInt(apiKeyID, 10) + liveKey := liveAPIKeySlotKeyPrefix + strconv.FormatInt(apiKeyID, 10) pipe.ZRemRangeByScore(ctx, slotKey, "-inf", strconv.FormatInt(cutoffTime, 10)) + pipe.ZRemRangeByScore(ctx, liveKey, "-inf", strconv.FormatInt(now.Unix()-liveLeaseTTLSeconds, 10)) cmds = append(cmds, apiKeyCmd{ apiKeyID: apiKeyID, zcardCmd: pipe.ZCard(ctx, slotKey), + liveCmd: pipe.ZCard(ctx, liveKey), }) } @@ -748,7 +879,7 @@ func (c *concurrencyCache) GetAPIKeyConcurrencyBatch(ctx context.Context, apiKey result := make(map[int64]int, len(apiKeyIDs)) for _, cmd := range cmds { - result[cmd.apiKeyID] = int(cmd.zcardCmd.Val()) + result[cmd.apiKeyID] = int(cmd.zcardCmd.Val() + cmd.liveCmd.Val()) } return result, nil } @@ -834,17 +965,21 @@ func (c *concurrencyCache) GetAccountsLoadBatch(ctx context.Context, accounts [] id int64 maxConcurrency int zcardCmd *redis.IntCmd + liveCmd *redis.IntCmd getCmd *redis.StringCmd } cmds := make([]accountCmds, 0, len(accounts)) for _, acc := range accounts { slotKey := accountSlotKeyPrefix + strconv.FormatInt(acc.ID, 10) + liveKey := liveAccountSlotKeyPrefix + strconv.FormatInt(acc.ID, 10) waitKey := accountWaitKeyPrefix + strconv.FormatInt(acc.ID, 10) pipe.ZRemRangeByScore(ctx, slotKey, "-inf", strconv.FormatInt(cutoffTime, 10)) + pipe.ZRemRangeByScore(ctx, liveKey, "-inf", strconv.FormatInt(now.Unix()-liveLeaseTTLSeconds, 10)) ac := accountCmds{ id: acc.ID, maxConcurrency: acc.MaxConcurrency, zcardCmd: pipe.ZCard(ctx, slotKey), + liveCmd: pipe.ZCard(ctx, liveKey), getCmd: pipe.Get(ctx, waitKey), } cmds = append(cmds, ac) @@ -856,7 +991,7 @@ func (c *concurrencyCache) GetAccountsLoadBatch(ctx context.Context, accounts [] loadMap := make(map[int64]*service.AccountLoadInfo, len(accounts)) for _, ac := range cmds { - currentConcurrency := int(ac.zcardCmd.Val()) + currentConcurrency := int(ac.zcardCmd.Val() + ac.liveCmd.Val()) waitingCount := 0 if v, err := ac.getCmd.Int(); err == nil { waitingCount = v @@ -894,17 +1029,21 @@ func (c *concurrencyCache) GetUsersLoadBatch(ctx context.Context, users []servic id int64 maxConcurrency int zcardCmd *redis.IntCmd + liveCmd *redis.IntCmd getCmd *redis.StringCmd } cmds := make([]userCmds, 0, len(users)) for _, u := range users { slotKey := userSlotKeyPrefix + strconv.FormatInt(u.ID, 10) + liveKey := liveUserSlotKeyPrefix + strconv.FormatInt(u.ID, 10) waitKey := waitQueueKeyPrefix + strconv.FormatInt(u.ID, 10) pipe.ZRemRangeByScore(ctx, slotKey, "-inf", strconv.FormatInt(cutoffTime, 10)) + pipe.ZRemRangeByScore(ctx, liveKey, "-inf", strconv.FormatInt(now.Unix()-liveLeaseTTLSeconds, 10)) uc := userCmds{ id: u.ID, maxConcurrency: u.MaxConcurrency, zcardCmd: pipe.ZCard(ctx, slotKey), + liveCmd: pipe.ZCard(ctx, liveKey), getCmd: pipe.Get(ctx, waitKey), } cmds = append(cmds, uc) @@ -916,7 +1055,7 @@ func (c *concurrencyCache) GetUsersLoadBatch(ctx context.Context, users []servic loadMap := make(map[int64]*service.UserLoadInfo, len(users)) for _, uc := range cmds { - currentConcurrency := int(uc.zcardCmd.Val()) + currentConcurrency := int(uc.zcardCmd.Val() + uc.liveCmd.Val()) waitingCount := 0 if v, err := uc.getCmd.Int(); err == nil { waitingCount = v diff --git a/backend/internal/repository/concurrency_cache_integration_test.go b/backend/internal/repository/concurrency_cache_integration_test.go index 02821159ee..4ec795c298 100644 --- a/backend/internal/repository/concurrency_cache_integration_test.go +++ b/backend/internal/repository/concurrency_cache_integration_test.go @@ -97,6 +97,42 @@ func (s *ConcurrencyCacheSuite) TestOpenAIWSIngressAPIKeySlot_ReapsCrashedLeaseW require.Equal(s.T(), int64(2), count) } +func (s *ConcurrencyCacheSuite) TestLiveLease_CountsTowardRegularAccountAndUserLimits() { + liveCache, ok := s.cache.(service.LiveConcurrencyCache) + require.True(s.T(), ok) + + accountID := int64(9101) + userID := int64(9102) + apiKeyID := int64(9103) + acquired, err := liveCache.AcquireLiveLease( + s.ctx, + accountID, + 1, + userID, + 1, + apiKeyID, + "live-integration", + false, + ) + require.NoError(s.T(), err) + require.True(s.T(), acquired) + + regularAccount, err := s.cache.AcquireAccountSlot(s.ctx, accountID, 1, "regular-account") + require.NoError(s.T(), err) + require.False(s.T(), regularAccount) + regularUser, err := s.cache.AcquireUserSlot(s.ctx, userID, 1, "regular-user") + require.NoError(s.T(), err) + require.False(s.T(), regularUser) + + refreshed, err := liveCache.RefreshLiveLease(s.ctx, accountID, userID, apiKeyID, "live-integration") + require.NoError(s.T(), err) + require.True(s.T(), refreshed) + require.NoError(s.T(), liveCache.ReleaseLiveLease(s.ctx, accountID, userID, apiKeyID, "live-integration")) + regularAccount, err = s.cache.AcquireAccountSlot(s.ctx, accountID, 1, "regular-account") + require.NoError(s.T(), err) + require.True(s.T(), regularAccount) +} + func (s *ConcurrencyCacheSuite) TestAccountSlot_AcquireAndRelease() { accountID := int64(10) reqID1, reqID2, reqID3 := "req1", "req2", "req3" diff --git a/backend/internal/repository/concurrency_cache_live_test.go b/backend/internal/repository/concurrency_cache_live_test.go new file mode 100644 index 0000000000..6dbebe50c4 --- /dev/null +++ b/backend/internal/repository/concurrency_cache_live_test.go @@ -0,0 +1,71 @@ +package repository + +import ( + "context" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" + "github.com/stretchr/testify/require" +) + +func TestLiveLeaseReplacesRegularSlotsAndCountsTowardLimits(t *testing.T) { + redisServer := miniredis.RunT(t) + client := redis.NewClient(&redis.Options{Addr: redisServer.Addr()}) + regular := NewConcurrencyCache(client, 15, 900) + live := regular.(service.LiveConcurrencyCache) + ctx := context.Background() + + accountAcquired, err := regular.AcquireAccountSlot(ctx, 10, 1, "regular-account") + require.NoError(t, err) + require.True(t, accountAcquired) + userAcquired, err := regular.AcquireUserSlot(ctx, 20, 1, "regular-user") + require.NoError(t, err) + require.True(t, userAcquired) + + acquired, err := live.AcquireLiveLease(ctx, 10, 1, 20, 1, 30, "live-lease", true) + require.NoError(t, err) + require.True(t, acquired) + require.NoError(t, regular.ReleaseAccountSlot(ctx, 10, "regular-account")) + require.NoError(t, regular.ReleaseUserSlot(ctx, 20, "regular-user")) + + accountCount, err := regular.GetAccountConcurrency(ctx, 10) + require.NoError(t, err) + require.Equal(t, 1, accountCount) + userCount, err := regular.GetUserConcurrency(ctx, 20) + require.NoError(t, err) + require.Equal(t, 1, userCount) + accountAcquired, err = regular.AcquireAccountSlot(ctx, 10, 1, "ordinary-blocked") + require.NoError(t, err) + require.False(t, accountAcquired) + + refreshed, err := live.RefreshLiveLease(ctx, 10, 20, 30, "live-lease") + require.NoError(t, err) + require.True(t, refreshed) + require.NoError(t, live.ReleaseLiveLease(ctx, 10, 20, 30, "live-lease")) + accountAcquired, err = regular.AcquireAccountSlot(ctx, 10, 1, "ordinary-allowed") + require.NoError(t, err) + require.True(t, accountAcquired) +} + +func TestLiveLeaseExpiresWithoutRefresh(t *testing.T) { + redisServer := miniredis.RunT(t) + client := redis.NewClient(&redis.Options{Addr: redisServer.Addr()}) + regular := NewConcurrencyCache(client, 15, 900) + live := regular.(service.LiveConcurrencyCache) + ctx := context.Background() + + acquired, err := live.AcquireLiveLease(ctx, 10, 1, 20, 1, 30, "expired-live", false) + require.NoError(t, err) + require.True(t, acquired) + + redisServer.FastForward(61 * time.Second) + acquired, err = regular.AcquireAccountSlot(ctx, 10, 1, "ordinary-after-expiry") + require.NoError(t, err) + require.True(t, acquired) + refreshed, err := live.RefreshLiveLease(ctx, 10, 20, 30, "expired-live") + require.NoError(t, err) + require.False(t, refreshed) +} diff --git a/backend/internal/repository/gateway_cache.go b/backend/internal/repository/gateway_cache.go index b8229eed1a..e17a135d52 100644 --- a/backend/internal/repository/gateway_cache.go +++ b/backend/internal/repository/gateway_cache.go @@ -2,7 +2,10 @@ package repository import ( "context" + "crypto/sha256" + "encoding/hex" "fmt" + "strconv" "time" "github.com/Wei-Shaw/sub2api/internal/service" @@ -10,6 +13,7 @@ import ( ) const stickySessionPrefix = "sticky_session:" +const liveCallPrefix = "live:call:" type gatewayCache struct { rdb *redis.Client @@ -54,6 +58,7 @@ func (c *gatewayCache) DeleteSessionAccountID(ctx context.Context, groupID int64 // Compile-time assertion: gatewayCache must implement CyberSessionBlockStore. var _ service.CyberSessionBlockStore = (*gatewayCache)(nil) +var _ service.LiveCallStore = (*gatewayCache)(nil) const cyberSessionBlockPrefix = "cyber_session_block:" @@ -71,3 +76,140 @@ func (c *gatewayCache) IsCyberSessionBlocked(ctx context.Context, key string) (b } return n > 0, nil } + +var claimLiveControllerScript = redis.NewScript(` + local key = KEYS[1] + local target = ARGV[1] + local owner = ARGV[2] + local current = redis.call('HGET', key, 'controller') + if current == false or current == 'closed' then + return 0 + end + if target == 'observer' and current ~= 'pending' then + return 0 + end + if target == 'proxy' and current ~= 'pending' and current ~= 'observer' and + (current ~= 'proxy' or redis.call('HGET', key, 'controller_owner') ~= owner) then + return 0 + end + redis.call('HSET', key, 'controller', target, 'controller_owner', owner) + return 1 +`) + +var markLiveCallClosedScript = redis.NewScript(` + local key = KEYS[1] + if redis.call('EXISTS', key) == 0 then + return 0 + end + if redis.call('HGET', key, 'controller') == 'closed' then + return 0 + end + redis.call('HSET', key, 'controller', 'closed', 'controller_owner', '') + redis.call('EXPIRE', key, ARGV[1]) + return 1 +`) + +var releaseLiveControllerScript = redis.NewScript(` + local key = KEYS[1] + if redis.call('HGET', key, 'controller') ~= 'proxy' or + redis.call('HGET', key, 'controller_owner') ~= ARGV[1] then + return 0 + end + redis.call('HSET', key, 'controller', 'pending', 'controller_owner', '') + return 1 +`) + +func liveCallKey(callHash string) string { + return liveCallPrefix + callHash +} + +func HashLiveCallID(callID string) string { + sum := sha256.Sum256([]byte(callID)) + return hex.EncodeToString(sum[:]) +} + +func (c *gatewayCache) SaveLiveCall(ctx context.Context, record *service.LiveCallRecord, ttl time.Duration) error { + if record == nil || record.CallHash == "" || record.CallID == "" { + return fmt.Errorf("invalid live call record") + } + values := map[string]any{ + "call_id": record.CallID, + "account_id": record.AccountID, + "api_key_id": record.APIKeyID, + "user_id": record.UserID, + "group_id": record.GroupID, + "subscription_id": record.SubscriptionID, + "lease_id": record.LeaseID, + "model": record.Model, + "created_at": record.CreatedAt.UnixMilli(), + "expires_at": record.ExpiresAt.UnixMilli(), + "controller": record.Controller, + "controller_owner": record.ControllerOwner, + "user_agent": record.UserAgent, + "ip_address": record.IPAddress, + "inbound_endpoint": record.InboundEndpoint, + } + key := liveCallKey(record.CallHash) + pipe := c.rdb.TxPipeline() + pipe.HSet(ctx, key, values) + pipe.Expire(ctx, key, ttl) + _, err := pipe.Exec(ctx) + return err +} + +func (c *gatewayCache) GetLiveCall(ctx context.Context, callHash string) (*service.LiveCallRecord, error) { + values, err := c.rdb.HGetAll(ctx, liveCallKey(callHash)).Result() + if err != nil { + return nil, err + } + if len(values) == 0 { + return nil, service.ErrLiveCallNotFound + } + parseInt := func(field string) int64 { + value, _ := strconv.ParseInt(values[field], 10, 64) + return value + } + createdAt := time.UnixMilli(parseInt("created_at")) + expiresAt := time.UnixMilli(parseInt("expires_at")) + return &service.LiveCallRecord{ + CallID: values["call_id"], + CallHash: callHash, + AccountID: parseInt("account_id"), + APIKeyID: parseInt("api_key_id"), + UserID: parseInt("user_id"), + GroupID: parseInt("group_id"), + SubscriptionID: parseInt("subscription_id"), + LeaseID: values["lease_id"], + Model: values["model"], + CreatedAt: createdAt, + ExpiresAt: expiresAt, + Controller: values["controller"], + ControllerOwner: values["controller_owner"], + UserAgent: values["user_agent"], + IPAddress: values["ip_address"], + InboundEndpoint: values["inbound_endpoint"], + }, nil +} + +func (c *gatewayCache) ClaimLiveController(ctx context.Context, callHash, controller, owner string) (bool, error) { + result, err := claimLiveControllerScript.Run(ctx, c.rdb, []string{liveCallKey(callHash)}, controller, owner).Int() + return result == 1, err +} + +func (c *gatewayCache) GetLiveController(ctx context.Context, callHash string) (string, error) { + value, err := c.rdb.HGet(ctx, liveCallKey(callHash), "controller").Result() + if err == redis.Nil { + return "", service.ErrLiveCallNotFound + } + return value, err +} + +func (c *gatewayCache) ReleaseLiveController(ctx context.Context, callHash, owner string) (bool, error) { + result, err := releaseLiveControllerScript.Run(ctx, c.rdb, []string{liveCallKey(callHash)}, owner).Int() + return result == 1, err +} + +func (c *gatewayCache) MarkLiveCallClosed(ctx context.Context, callHash string, ttl time.Duration) (bool, error) { + result, err := markLiveCallClosedScript.Run(ctx, c.rdb, []string{liveCallKey(callHash)}, int64(ttl.Seconds())).Int() + return result == 1, err +} diff --git a/backend/internal/repository/gateway_cache_live_test.go b/backend/internal/repository/gateway_cache_live_test.go new file mode 100644 index 0000000000..ad2b4e804b --- /dev/null +++ b/backend/internal/repository/gateway_cache_live_test.go @@ -0,0 +1,58 @@ +package repository + +import ( + "context" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" + "github.com/stretchr/testify/require" +) + +func TestGatewayCacheLiveCallIdentityAndController(t *testing.T) { + redisServer := miniredis.RunT(t) + client := redis.NewClient(&redis.Options{Addr: redisServer.Addr()}) + cache := NewGatewayCache(client).(service.LiveCallStore) + otherInstance := NewGatewayCache(client).(service.LiveCallStore) + record := &service.LiveCallRecord{ + CallID: "call_secret", + CallHash: HashLiveCallID("call_secret"), + AccountID: 11, + APIKeyID: 22, + UserID: 33, + GroupID: 44, + LeaseID: "lease", + Model: "gpt-live-test", + CreatedAt: time.Now(), + ExpiresAt: time.Now().Add(time.Hour), + Controller: service.LiveControllerPending, + } + require.NoError(t, cache.SaveLiveCall(context.Background(), record, time.Hour)) + + loaded, err := otherInstance.GetLiveCall(context.Background(), record.CallHash) + require.NoError(t, err) + require.Equal(t, record.CallID, loaded.CallID) + require.Equal(t, record.AccountID, loaded.AccountID) + + claimed, err := cache.ClaimLiveController(context.Background(), record.CallHash, service.LiveControllerObserver, "observer-1") + require.NoError(t, err) + require.True(t, claimed) + claimed, err = cache.ClaimLiveController(context.Background(), record.CallHash, service.LiveControllerProxy, "proxy-1") + require.NoError(t, err) + require.True(t, claimed) + controller, err := cache.GetLiveController(context.Background(), record.CallHash) + require.NoError(t, err) + require.Equal(t, service.LiveControllerProxy, controller) + + released, err := cache.ReleaseLiveController(context.Background(), record.CallHash, "proxy-1") + require.NoError(t, err) + require.True(t, released) + closed, err := cache.MarkLiveCallClosed(context.Background(), record.CallHash, time.Hour) + require.NoError(t, err) + require.True(t, closed) + closed, err = cache.MarkLiveCallClosed(context.Background(), record.CallHash, time.Hour) + require.NoError(t, err) + require.False(t, closed) +} diff --git a/backend/internal/repository/group_repo.go b/backend/internal/repository/group_repo.go index 96be19a70c..50c9aafa5f 100644 --- a/backend/internal/repository/group_repo.go +++ b/backend/internal/repository/group_repo.go @@ -90,6 +90,7 @@ func createGroupRecord(ctx context.Context, client *dbent.Client, groupIn *servi SetModelRoutingEnabled(groupIn.ModelRoutingEnabled). SetMcpXMLInject(groupIn.MCPXMLInject). SetAllowMessagesDispatch(groupIn.AllowMessagesDispatch). + SetAllowLive(groupIn.AllowLive). SetRequireOauthOnly(groupIn.RequireOAuthOnly). SetRequirePrivacySet(groupIn.RequirePrivacySet). SetDefaultMappedModel(groupIn.DefaultMappedModel). @@ -255,6 +256,7 @@ func (r *groupRepository) Update(ctx context.Context, groupIn *service.Group) er SetModelRoutingEnabled(groupIn.ModelRoutingEnabled). SetMcpXMLInject(groupIn.MCPXMLInject). SetAllowMessagesDispatch(groupIn.AllowMessagesDispatch). + SetAllowLive(groupIn.AllowLive). SetRequireOauthOnly(groupIn.RequireOAuthOnly). SetRequirePrivacySet(groupIn.RequirePrivacySet). SetDefaultMappedModel(groupIn.DefaultMappedModel). diff --git a/backend/internal/repository/migrations_schema_integration_test.go b/backend/internal/repository/migrations_schema_integration_test.go index 68eb790edb..b393bc15b6 100644 --- a/backend/internal/repository/migrations_schema_integration_test.go +++ b/backend/internal/repository/migrations_schema_integration_test.go @@ -55,6 +55,9 @@ func TestMigrationsRunner_IsIdempotent_AndSchemaIsUpToDate(t *testing.T) { requireColumn(t, tx, "accounts", "session_window_status", "character varying", 20, true) requireIndex(t, tx, "accounts", "idx_accounts_autopause_expiry_due") + // groups: OpenAI Live 默认关闭,管理员显式开启后才可访问。 + requireColumn(t, tx, "groups", "allow_live", "boolean", 0, false) + // api_keys: key length should be 128 requireColumn(t, tx, "api_keys", "key", "character varying", 128, false) diff --git a/backend/internal/server/api_contract_test.go b/backend/internal/server/api_contract_test.go index 11b8eb5f76..7ab057ef0a 100644 --- a/backend/internal/server/api_contract_test.go +++ b/backend/internal/server/api_contract_test.go @@ -377,6 +377,7 @@ func TestAPIContracts(t *testing.T) { "video_rate_multiplier": 0, "claude_code_only": false, "allow_messages_dispatch": false, + "allow_live": false, "fallback_group_id": null, "fallback_group_id_on_invalid_request": null, "require_oauth_only": false, diff --git a/backend/internal/server/routes/gateway.go b/backend/internal/server/routes/gateway.go index ee435eca8c..b059646fad 100644 --- a/backend/internal/server/routes/gateway.go +++ b/backend/internal/server/routes/gateway.go @@ -180,6 +180,8 @@ func RegisterGatewayRoutes( // Codex manifest format; other clients keep the OpenAI-style list. gateway.GET("/models", modelsHandler) gateway.GET("/usage", h.Gateway.Usage) + gateway.POST("/live", h.OpenAIGateway.Live) + gateway.GET("/live/:call_id", h.OpenAIGateway.LiveSideband) // OpenAI Responses API: auto-route based on group platform gateway.POST("/responses", func(c *gin.Context) { if isOpenAIResponsesCompatibleGatewayPlatform(c) { @@ -277,6 +279,8 @@ func RegisterGatewayRoutes( codexDirect := r.Group("/backend-api/codex") codexDirect.Use(bodyLimit, clientRequestID, opsErrorLogger, endpointNorm, gin.HandlerFunc(apiKeyAuth), compositeTarget, requireGroupAnthropic) { + codexDirect.POST("/realtime/calls", h.OpenAIGateway.Live) + codexDirect.GET("/:call_id", h.OpenAIGateway.LiveSideband) codexDirect.POST("/responses", responsesHandler) codexDirect.POST("/responses/*subpath", responsesHandler) codexDirect.POST("/alpha/search", textBodyLimit, h.OpenAIGateway.AlphaSearch) diff --git a/backend/internal/server/routes/prompt_audit_route_coverage_test.go b/backend/internal/server/routes/prompt_audit_route_coverage_test.go index ba38910b4b..1620a18318 100644 --- a/backend/internal/server/routes/prompt_audit_route_coverage_test.go +++ b/backend/internal/server/routes/prompt_audit_route_coverage_test.go @@ -34,6 +34,8 @@ func TestEveryGatewayPOSTRouteIsClassifiedForPromptAuditCoverage(t *testing.T) { "/chat/completions": {"gateway_handler_chat_completions.go", "openai_chat_completions.go"}, "/embeddings": {"openai_embeddings.go"}, "/alpha/search": {"openai_alpha_search.go"}, + "/live": {"openai_live.go"}, + "/realtime/calls": {"openai_live.go"}, "/images/generations": {"openai_images.go", "grok_media.go"}, "/images/edits": {"openai_images.go", "grok_media.go"}, "/images/generations/async": {"image_task_handler.go"}, diff --git a/backend/internal/service/account.go b/backend/internal/service/account.go index 5b3707d2a0..b7f6d2266a 100644 --- a/backend/internal/service/account.go +++ b/backend/internal/service/account.go @@ -89,6 +89,7 @@ const ( OpenAIEndpointCapabilityChatCompletions OpenAIEndpointCapability = "chat_completions" OpenAIEndpointCapabilityEmbeddings OpenAIEndpointCapability = "embeddings" OpenAIEndpointCapabilityAlphaSearch OpenAIEndpointCapability = "alpha_search" + OpenAIEndpointCapabilityLive OpenAIEndpointCapability = "live" // OpenAIEndpointCapabilityGrokMediaGeneration keeps image/video generation // away from Grok accounts that are explicitly disabled or whose billing // entitlement probe was forbidden. Video status lookups intentionally do not @@ -1435,6 +1436,11 @@ func (a *Account) SupportsOpenAIEndpointCapability(capability OpenAIEndpointCapa } switch capability { case OpenAIEndpointCapabilityChatCompletions: + case OpenAIEndpointCapabilityLive: + return a.Platform == PlatformOpenAI && + a.Type == AccountTypeOAuth && + !a.IsOpenAIPersonalAccessToken() && + !a.IsOpenAIAgentIdentity() case OpenAIEndpointCapabilityResponses: // Responses 支持状态由 accounts.extra 的自动探测标记决定,而非 // credentials 能力集。已探测确认不支持 /v1/responses 的 APIKey 上游 diff --git a/backend/internal/service/admin_group.go b/backend/internal/service/admin_group.go index fdba370ea0..b773c43ed3 100644 --- a/backend/internal/service/admin_group.go +++ b/backend/internal/service/admin_group.go @@ -470,6 +470,7 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn MCPXMLInject: mcpXMLInject, SupportedModelScopes: input.SupportedModelScopes, AllowMessagesDispatch: input.AllowMessagesDispatch, + AllowLive: input.AllowLive, RequireOAuthOnly: input.RequireOAuthOnly, RequirePrivacySet: input.RequirePrivacySet, DefaultMappedModel: input.DefaultMappedModel, @@ -480,6 +481,9 @@ func (s *adminServiceImpl) CreateGroup(ctx context.Context, input *CreateGroupIn ReasoningEffortMappings: reasoningEffortMappings, } sanitizeGroupMessagesDispatchFields(group) + if group.Platform != PlatformOpenAI { + group.AllowLive = false + } sanitizeGroupReasoningEffortPolicy(group) if err := s.groupRepo.Create(ctx, group); err != nil { return nil, err @@ -777,6 +781,9 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd if input.AllowMessagesDispatch != nil { group.AllowMessagesDispatch = *input.AllowMessagesDispatch } + if input.AllowLive != nil { + group.AllowLive = *input.AllowLive + } if input.RequireOAuthOnly != nil { group.RequireOAuthOnly = *input.RequireOAuthOnly } @@ -810,6 +817,9 @@ func (s *adminServiceImpl) UpdateGroup(ctx context.Context, id int64, input *Upd group.ReasoningEffortMappings = reasoningEffortMappings } sanitizeGroupMessagesDispatchFields(group) + if group.Platform != PlatformOpenAI { + group.AllowLive = false + } sanitizeGroupReasoningEffortPolicy(group) if err := s.groupRepo.Update(ctx, group); err != nil { diff --git a/backend/internal/service/admin_group_duplicate.go b/backend/internal/service/admin_group_duplicate.go index 6c57df1cf0..fcfe39326d 100644 --- a/backend/internal/service/admin_group_duplicate.go +++ b/backend/internal/service/admin_group_duplicate.go @@ -120,6 +120,7 @@ func cloneGroupForDuplicate(source *Group, operationID string) *Group { SupportedModelScopes: append([]string(nil), source.SupportedModelScopes...), SortOrder: source.SortOrder, AllowMessagesDispatch: source.AllowMessagesDispatch, + AllowLive: source.AllowLive, RequireOAuthOnly: source.RequireOAuthOnly, RequirePrivacySet: source.RequirePrivacySet, DefaultMappedModel: source.DefaultMappedModel, diff --git a/backend/internal/service/admin_group_duplicate_test.go b/backend/internal/service/admin_group_duplicate_test.go index c6c48d1e1f..44db1afdc7 100644 --- a/backend/internal/service/admin_group_duplicate_test.go +++ b/backend/internal/service/admin_group_duplicate_test.go @@ -162,6 +162,7 @@ func TestDuplicateGroupCopiesConfigurationDeeplyAndResetsRuntimeState(t *testing SupportedModelScopes: []string{"claude", "gemini_text"}, SortOrder: 9, AllowMessagesDispatch: true, + AllowLive: true, RequireOAuthOnly: true, RequirePrivacySet: true, DefaultMappedModel: "gpt-5.4", diff --git a/backend/internal/service/admin_service.go b/backend/internal/service/admin_service.go index 49c5272ed0..bd907125a5 100644 --- a/backend/internal/service/admin_service.go +++ b/backend/internal/service/admin_service.go @@ -252,6 +252,7 @@ type CreateGroupInput struct { SupportedModelScopes []string // OpenAI Messages 调度配置(仅 openai 平台使用) AllowMessagesDispatch bool + AllowLive bool DefaultMappedModel string RequireOAuthOnly bool RequirePrivacySet bool @@ -312,6 +313,7 @@ type UpdateGroupInput struct { SupportedModelScopes *[]string // OpenAI Messages 调度配置(仅 openai 平台使用) AllowMessagesDispatch *bool + AllowLive *bool DefaultMappedModel *string RequireOAuthOnly *bool RequirePrivacySet *bool diff --git a/backend/internal/service/admin_service_group_test.go b/backend/internal/service/admin_service_group_test.go index cf13417879..98e63899b6 100644 --- a/backend/internal/service/admin_service_group_test.go +++ b/backend/internal/service/admin_service_group_test.go @@ -957,6 +957,7 @@ func TestAdminService_CreateGroup_ClearsMessagesDispatchFieldsForNonOpenAIPlatfo Platform: PlatformAnthropic, RateMultiplier: 1.0, AllowMessagesDispatch: true, + AllowLive: true, DefaultMappedModel: "gpt-5.4", MessagesDispatchModelConfig: OpenAIMessagesDispatchModelConfig{ OpusMappedModel: "gpt-5.4", @@ -966,6 +967,7 @@ func TestAdminService_CreateGroup_ClearsMessagesDispatchFieldsForNonOpenAIPlatfo require.NotNil(t, group) require.NotNil(t, repo.created) require.False(t, repo.created.AllowMessagesDispatch) + require.False(t, repo.created.AllowLive) require.Empty(t, repo.created.DefaultMappedModel) require.Equal(t, OpenAIMessagesDispatchModelConfig{}, repo.created.MessagesDispatchModelConfig) } @@ -977,6 +979,7 @@ func TestAdminService_UpdateGroup_ClearsMessagesDispatchFieldsWhenPlatformChange Platform: PlatformOpenAI, Status: StatusActive, AllowMessagesDispatch: true, + AllowLive: true, DefaultMappedModel: "gpt-5.4", MessagesDispatchModelConfig: OpenAIMessagesDispatchModelConfig{ SonnetMappedModel: "gpt-5.3-codex", @@ -993,6 +996,7 @@ func TestAdminService_UpdateGroup_ClearsMessagesDispatchFieldsWhenPlatformChange require.NotNil(t, repo.updated) require.Equal(t, PlatformAnthropic, repo.updated.Platform) require.False(t, repo.updated.AllowMessagesDispatch) + require.False(t, repo.updated.AllowLive) require.Empty(t, repo.updated.DefaultMappedModel) require.Equal(t, OpenAIMessagesDispatchModelConfig{}, repo.updated.MessagesDispatchModelConfig) } diff --git a/backend/internal/service/api_key_auth_cache.go b/backend/internal/service/api_key_auth_cache.go index 94799de52d..924fc20fcd 100644 --- a/backend/internal/service/api_key_auth_cache.go +++ b/backend/internal/service/api_key_auth_cache.go @@ -94,6 +94,7 @@ type APIKeyAuthGroupSnapshot struct { // OpenAI Messages 调度配置(仅 openai 平台使用) AllowMessagesDispatch bool `json:"allow_messages_dispatch"` + AllowLive bool `json:"allow_live"` DefaultMappedModel string `json:"default_mapped_model,omitempty"` MessagesDispatchModelConfig OpenAIMessagesDispatchModelConfig `json:"messages_dispatch_model_config,omitempty"` ModelsListConfig GroupModelsListConfig `json:"models_list_config,omitempty"` diff --git a/backend/internal/service/api_key_auth_cache_impl.go b/backend/internal/service/api_key_auth_cache_impl.go index 94c4276bf1..fd5fda22c7 100644 --- a/backend/internal/service/api_key_auth_cache_impl.go +++ b/backend/internal/service/api_key_auth_cache_impl.go @@ -14,7 +14,7 @@ import ( "github.com/dgraph-io/ristretto" ) -const apiKeyAuthSnapshotVersion = 16 // v16: include group reasoning effort ceiling and mappings +const apiKeyAuthSnapshotVersion = 17 // v17: include the OpenAI group Live gate type apiKeyAuthCacheConfig struct { l1Size int @@ -409,6 +409,7 @@ func (s *APIKeyService) snapshotFromAPIKey(ctx context.Context, apiKey *APIKey) MCPXMLInject: apiKey.Group.MCPXMLInject, SupportedModelScopes: apiKey.Group.SupportedModelScopes, AllowMessagesDispatch: apiKey.Group.AllowMessagesDispatch, + AllowLive: apiKey.Group.AllowLive, DefaultMappedModel: apiKey.Group.DefaultMappedModel, MessagesDispatchModelConfig: apiKey.Group.MessagesDispatchModelConfig, ModelsListConfig: apiKey.Group.ModelsListConfig, @@ -495,6 +496,7 @@ func (s *APIKeyService) snapshotToAPIKey(key string, snapshot *APIKeyAuthSnapsho MCPXMLInject: snapshot.Group.MCPXMLInject, SupportedModelScopes: snapshot.Group.SupportedModelScopes, AllowMessagesDispatch: snapshot.Group.AllowMessagesDispatch, + AllowLive: snapshot.Group.AllowLive, DefaultMappedModel: snapshot.Group.DefaultMappedModel, MessagesDispatchModelConfig: snapshot.Group.MessagesDispatchModelConfig, ModelsListConfig: snapshot.Group.ModelsListConfig, diff --git a/backend/internal/service/group.go b/backend/internal/service/group.go index e428035db5..b71e6b9e84 100644 --- a/backend/internal/service/group.go +++ b/backend/internal/service/group.go @@ -82,6 +82,7 @@ type Group struct { // OpenAI Messages 调度配置(仅 openai 平台使用) AllowMessagesDispatch bool + AllowLive bool RequireOAuthOnly bool // 仅允许非 apikey 类型账号关联(OpenAI/Antigravity/Anthropic/Gemini) RequirePrivacySet bool // 调度时仅允许 privacy 已成功设置的账号(OpenAI/Antigravity/Anthropic/Gemini) DefaultMappedModel string diff --git a/backend/internal/service/openai_live.go b/backend/internal/service/openai_live.go new file mode 100644 index 0000000000..9ec872ed31 --- /dev/null +++ b/backend/internal/service/openai_live.go @@ -0,0 +1,747 @@ +package service + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/url" + "path" + "strings" + "time" + + "github.com/Wei-Shaw/sub2api/internal/pkg/logger" + coderws "github.com/coder/websocket" + "github.com/google/uuid" + "github.com/tidwall/gjson" + "go.uber.org/zap" +) + +const ( + defaultLiveMaxSessionDuration = time.Hour + liveLeaseRefreshInterval = 20 * time.Second + liveRedisOperationTimeout = 3 * time.Second + liveClosedRecordTTL = 24 * time.Hour + liveObserverPollInterval = 250 * time.Millisecond + liveUpstreamBodyLimit = 2 << 20 +) + +var ( + chatGPTLiveCallsURL = "https://chatgpt.com/backend-api/codex/realtime/calls?intent=quicksilver&architecture=avas" + chatGPTLiveSidebandBaseURL = "wss://chatgpt.com/backend-api/codex" +) + +type liveFrameConn interface { + ReadFrame(ctx context.Context) (coderws.MessageType, []byte, error) + WriteFrame(ctx context.Context, msgType coderws.MessageType, payload []byte) error + Close() error +} + +func liveSidebandReadError(err error) error { + if coderws.CloseStatus(err) == coderws.StatusNormalClosure { + return ErrLiveCallNotFound + } + return err +} + +func hashLiveCallID(callID string) string { + sum := sha256.Sum256([]byte(callID)) + return hex.EncodeToString(sum[:]) +} + +func liveGroupID(groupID *int64) int64 { + if groupID == nil { + return 0 + } + return *groupID +} + +func liveOptionalID(value int64) *int64 { + if value <= 0 { + return nil + } + result := value + return &result +} + +func (s *OpenAIGatewayService) liveStore() (LiveCallStore, error) { + if s == nil || s.cache == nil { + return nil, ErrLiveUnavailable + } + store, ok := s.cache.(LiveCallStore) + if !ok { + return nil, ErrLiveUnavailable + } + return store, nil +} + +func (s *OpenAIGatewayService) liveConcurrencyCache() (LiveConcurrencyCache, error) { + if s == nil || s.concurrencyService == nil || s.concurrencyService.cache == nil { + return nil, ErrLiveUnavailable + } + cache, ok := s.concurrencyService.cache.(LiveConcurrencyCache) + if !ok { + return nil, ErrLiveUnavailable + } + return cache, nil +} + +func (s *OpenAIGatewayService) liveMaxSessionDuration() time.Duration { + if s != nil && s.cfg != nil && s.cfg.Gateway.Live.MaxSessionDurationSeconds > 0 { + return time.Duration(s.cfg.Gateway.Live.MaxSessionDurationSeconds) * time.Second + } + return defaultLiveMaxSessionDuration +} + +func ValidateLiveCallRequest(request *LiveCallRequest) error { + if request == nil || strings.TrimSpace(request.SDP) == "" { + return errors.New("sdp is required") + } + if len(request.Session) == 0 || !json.Valid(request.Session) { + return errors.New("session must be valid JSON") + } + var sessionObject map[string]json.RawMessage + if err := json.Unmarshal(request.Session, &sessionObject); err != nil { + return errors.New("session must be a JSON object") + } + if sessionObject == nil { + return errors.New("session must be a JSON object") + } + return nil +} + +// CreateLiveCall 创建 Frameless 会话。调用方须在调用期间持有普通用户槽位; +// 调度器持有的普通账号槽位会被同一个 Live 租约原子接替。 +func (s *OpenAIGatewayService) CreateLiveCall( + ctx context.Context, + request *LiveCallRequest, + identity LiveCallIdentity, + userMaxConcurrency int, +) (*LiveCallCreated, error) { + if err := ValidateLiveCallRequest(request); err != nil { + return nil, err + } + store, err := s.liveStore() + if err != nil { + return nil, err + } + liveCache, err := s.liveConcurrencyCache() + if err != nil { + return nil, err + } + + excluded := make(map[int64]struct{}) + var lastErr error + for attempt := 0; attempt <= 3; attempt++ { + selection, _, selectErr := s.SelectAccountWithSchedulerForCapability( + ctx, + identity.GroupID, + "", + uuid.NewString(), + "", + excluded, + OpenAIUpstreamTransportHTTPSSE, + OpenAIEndpointCapabilityLive, + false, + false, + false, + ) + if selectErr != nil { + if lastErr != nil { + return nil, lastErr + } + return nil, selectErr + } + if selection == nil || selection.Account == nil || !selection.Acquired { + if selection != nil && selection.ReleaseFunc != nil { + selection.ReleaseFunc() + } + return nil, ErrLiveConcurrencyFull + } + + account := selection.Account + leaseID := generateRequestID() + acquired, acquireErr := liveCache.AcquireLiveLease( + ctx, + account.ID, + account.Concurrency, + identity.UserID, + userMaxConcurrency, + identity.APIKeyID, + leaseID, + true, + ) + if acquireErr != nil || !acquired { + selection.ReleaseFunc() + if acquireErr != nil { + return nil, acquireErr + } + return nil, ErrLiveConcurrencyFull + } + + created, createErr := s.createUpstreamLiveCall(ctx, account, request) + selection.ReleaseFunc() + if createErr != nil { + s.releaseLiveLease(account.ID, identity.UserID, identity.APIKeyID, leaseID) + if !s.shouldFailoverLiveCreateError(createErr) { + return nil, createErr + } + excluded[account.ID] = struct{}{} + lastErr = createErr + continue + } + + now := time.Now() + model := strings.TrimSpace(gjson.GetBytes(request.Session, "model").String()) + if model == "" { + model = "gpt-live" + } + record := &LiveCallRecord{ + CallID: created.CallID, + CallHash: hashLiveCallID(created.CallID), + AccountID: account.ID, + APIKeyID: identity.APIKeyID, + UserID: identity.UserID, + GroupID: liveGroupID(identity.GroupID), + SubscriptionID: liveGroupID(identity.SubscriptionID), + LeaseID: leaseID, + Model: model, + CreatedAt: now, + ExpiresAt: now.Add(s.liveMaxSessionDuration()), + Controller: LiveControllerPending, + UserAgent: identity.UserAgent, + IPAddress: identity.IPAddress, + InboundEndpoint: identity.InboundEndpoint, + } + mappingTTL := s.liveMaxSessionDuration() + 5*time.Minute + if saveErr := store.SaveLiveCall(ctx, record, mappingTTL); saveErr != nil { + s.releaseLiveLease(account.ID, identity.UserID, identity.APIKeyID, leaseID) + return nil, fmt.Errorf("save live call mapping: %w", saveErr) + } + created.Account = account + go s.observeLiveCall(record.CallHash) + return created, nil + } + if lastErr != nil { + return nil, lastErr + } + return nil, ErrLiveUnavailable +} + +func (s *OpenAIGatewayService) shouldFailoverLiveCreateError(err error) bool { + var upstreamErr *UpstreamFailoverError + if !errors.As(err, &upstreamErr) { + // 凭证读取和网络传输错误都可能只影响当前账号或代理。 + return true + } + return s.shouldFailoverOpenAIUpstreamResponse( + upstreamErr.StatusCode, + "", + upstreamErr.ResponseBody, + ) +} + +func (s *OpenAIGatewayService) createUpstreamLiveCall( + ctx context.Context, + account *Account, + request *LiveCallRequest, +) (*LiveCallCreated, error) { + token, _, err := s.GetAccessToken(ctx, account) + if err != nil { + return nil, err + } + body, err := json.Marshal(struct { + SDP string `json:"sdp"` + Session json.RawMessage `json:"session"` + }{ + SDP: request.SDP, + Session: request.Session, + }) + if err != nil { + return nil, err + } + reqCtx := WithHTTPUpstreamRedirectsDisabled(WithHTTPUpstreamProfile(ctx, HTTPUpstreamProfileOpenAI)) + upstreamReq, err := http.NewRequestWithContext(reqCtx, http.MethodPost, chatGPTLiveCallsURL, bytes.NewReader(body)) + if err != nil { + return nil, err + } + authHeaders, err := s.buildOpenAIAuthenticationHeaders(ctx, account, token) + if err != nil { + return nil, err + } + for key, values := range authHeaders { + for _, value := range values { + upstreamReq.Header.Add(key, value) + } + } + upstreamReq.Host = "chatgpt.com" + if err := resolveAndSetOpenAIChatGPTAccountHeaders(ctx, s.accountRepo, upstreamReq.Header, account); err != nil { + return nil, err + } + upstreamReq.Header.Set("Content-Type", "application/json") + upstreamReq.Header.Set("Accept", "application/sdp") + applyLiveUpstreamIdentityHeaders(upstreamReq.Header) + + resp, err := s.httpUpstream.Do(upstreamReq, resolveAccountProxyURL(account), account.ID, account.Concurrency) + if err != nil { + return nil, err + } + defer func() { _ = resp.Body.Close() }() + responseBody, readErr := io.ReadAll(io.LimitReader(resp.Body, liveUpstreamBodyLimit+1)) + if readErr != nil { + return nil, readErr + } + if len(responseBody) > liveUpstreamBodyLimit { + return nil, errors.New("live upstream response is too large") + } + if resp.StatusCode < 200 || resp.StatusCode >= 300 { + logLiveUpstreamFailure(ctx, account.ID, resp.StatusCode, resp.Header, responseBody) + return nil, &UpstreamFailoverError{ + StatusCode: resp.StatusCode, + ResponseBody: responseBody, + ResponseHeaders: resp.Header.Clone(), + } + } + callID, err := liveCallIDFromLocation(resp.Header.Get("Location")) + if err != nil { + return nil, err + } + return &LiveCallCreated{ + SDP: responseBody, + CallID: callID, + Location: resp.Header.Get("Location"), + }, nil +} + +func logLiveUpstreamFailure( + ctx context.Context, + accountID int64, + statusCode int, + headers http.Header, + body []byte, +) { + errorType := strings.TrimSpace(gjson.GetBytes(body, "error.type").String()) + errorCode := strings.TrimSpace(gjson.GetBytes(body, "error.code").String()) + errorMessage := strings.TrimSpace(gjson.GetBytes(body, "error.message").String()) + if errorType == "" { + errorType = strings.TrimSpace(gjson.GetBytes(body, "type").String()) + } + if errorCode == "" { + errorCode = strings.TrimSpace(gjson.GetBytes(body, "code").String()) + } + if errorMessage == "" { + errorMessage = strings.TrimSpace(gjson.GetBytes(body, "message").String()) + } + if errorMessage == "" { + errorMessage = strings.TrimSpace(gjson.GetBytes(body, "detail").String()) + } + + logger.FromContext(ctx).Warn( + "OpenAI Live 上游拒绝请求", + zap.Int64("account_id", accountID), + zap.Int("upstream_status_code", statusCode), + zap.String("upstream_error_type", truncateOpenAIWSLogValue(errorType, 120)), + zap.String("upstream_error_code", truncateOpenAIWSLogValue(errorCode, 120)), + zap.String("upstream_error_message", truncateOpenAIWSLogValue(errorMessage, 300)), + zap.String("upstream_content_type", truncateOpenAIWSLogValue(headers.Get("Content-Type"), 120)), + zap.String("upstream_server", truncateOpenAIWSLogValue(headers.Get("Server"), 120)), + zap.String("upstream_cf_mitigated", truncateOpenAIWSLogValue(headers.Get("Cf-Mitigated"), 120)), + zap.String("upstream_cf_ray", truncateOpenAIWSLogValue(headers.Get("Cf-Ray"), 120)), + zap.String("upstream_request_id", truncateOpenAIWSLogValue(headers.Get("X-Request-Id"), 120)), + ) +} + +func liveCallIDFromLocation(location string) (string, error) { + location = strings.TrimSpace(location) + if location == "" { + return "", errors.New("live upstream response has no Location") + } + parsed, err := url.Parse(location) + if err != nil { + return "", fmt.Errorf("parse live Location: %w", err) + } + callID := strings.TrimSpace(path.Base(strings.TrimSuffix(parsed.Path, "/"))) + if callID == "" || callID == "." || callID == "codex" { + return "", errors.New("live upstream Location has no call id") + } + return callID, nil +} + +func applyLiveUpstreamIdentityHeaders(headers http.Header) { + headers.Set("OpenAI-Alpha", "quicksilver=v2") + ensureCodexIdentityHeaders(headers) + enforceCodexIdentityHeaders(headers) + // Realtime/Live 不使用 Responses 的实验头。 + headers.Del("OpenAI-Beta") +} + +func (s *OpenAIGatewayService) liveSidebandHeaders(ctx context.Context, account *Account) (http.Header, error) { + token, _, err := s.GetAccessToken(ctx, account) + if err != nil { + return nil, err + } + headers, err := s.buildOpenAIAuthenticationHeaders(ctx, account, token) + if err != nil { + return nil, err + } + if err := resolveAndSetOpenAIChatGPTAccountHeaders(ctx, s.accountRepo, headers, account); err != nil { + return nil, err + } + applyLiveUpstreamIdentityHeaders(headers) + return headers, nil +} + +func (s *OpenAIGatewayService) dialLiveSideband(ctx context.Context, record *LiveCallRecord) (liveFrameConn, error) { + account, err := s.accountRepo.GetByID(ctx, record.AccountID) + if err != nil { + return nil, err + } + if account == nil || !account.SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityLive) { + return nil, ErrLiveUnavailable + } + headers, err := s.liveSidebandHeaders(ctx, account) + if err != nil { + return nil, err + } + target := strings.TrimRight(chatGPTLiveSidebandBaseURL, "/") + "/" + url.PathEscape(record.CallID) + conn, status, _, err := s.getOpenAIWSPassthroughDialer().Dial(ctx, target, headers, resolveAccountProxyURL(account)) + if err != nil { + return nil, fmt.Errorf("dial live sideband (status %d): %w", status, err) + } + raw, ok := conn.(liveFrameConn) + if !ok { + _ = conn.Close() + return nil, errors.New("live sideband transport does not support raw frames") + } + return raw, nil +} + +func (s *OpenAIGatewayService) GetLiveCallForIdentity( + ctx context.Context, + callID string, + identity LiveCallIdentity, +) (*LiveCallRecord, error) { + store, err := s.liveStore() + if err != nil { + return nil, err + } + record, err := store.GetLiveCall(ctx, hashLiveCallID(callID)) + if err != nil { + return nil, err + } + if record.CallID != callID || + record.APIKeyID != identity.APIKeyID || + record.UserID != identity.UserID || + record.GroupID != liveGroupID(identity.GroupID) { + return nil, ErrLiveIdentityMismatch + } + if record.Controller == LiveControllerClosed { + return nil, ErrLiveCallNotFound + } + return record, nil +} + +// ProxyLiveSideband 让认证后的客户端接管控制连接;媒体始终不经过这里。 +func (s *OpenAIGatewayService) ProxyLiveSideband( + ctx context.Context, + record *LiveCallRecord, + downstream *coderws.Conn, +) error { + if record == nil || downstream == nil { + return ErrLiveCallNotFound + } + store, err := s.liveStore() + if err != nil { + return err + } + owner := uuid.NewString() + claimed, err := store.ClaimLiveController(ctx, record.CallHash, LiveControllerProxy, owner) + if err != nil { + return err + } + if !claimed { + return ErrLiveControllerChanged + } + + // observer 轮询到接管状态后会关闭旧控制连接;同一个 call 可重新加入。 + time.Sleep(liveObserverPollInterval) + upstream, err := s.dialLiveSideband(ctx, record) + if err != nil { + _, _ = store.ReleaseLiveController(context.Background(), record.CallHash, owner) + go s.observeLiveCall(record.CallHash) + return err + } + defer upstream.Close() + downstream.SetReadLimit(openAIWSMessageReadLimitBytes) + + proxyCtx, cancel := context.WithCancel(ctx) + defer cancel() + errCh := make(chan error, 2) + go func() { + for { + messageType, payload, readErr := downstream.Read(proxyCtx) + if readErr != nil { + errCh <- readErr + return + } + if writeErr := upstream.WriteFrame(proxyCtx, messageType, payload); writeErr != nil { + errCh <- writeErr + return + } + } + }() + go func() { + for { + messageType, payload, readErr := upstream.ReadFrame(proxyCtx) + if readErr != nil { + errCh <- liveSidebandReadError(readErr) + return + } + if writeErr := downstream.Write(proxyCtx, messageType, payload); writeErr != nil { + errCh <- writeErr + return + } + if messageType == coderws.MessageText { + eventType := strings.TrimSpace(gjson.GetBytes(payload, "type").String()) + if eventType == "session.closed" || eventType == "session.ended" { + errCh <- ErrLiveCallNotFound + return + } + } + } + }() + + runErr := s.runLiveController(proxyCtx, record, upstream, errCh) + cancel() + _, _ = store.ReleaseLiveController(context.Background(), record.CallHash, owner) + if errors.Is(runErr, ErrLiveCallNotFound) { + s.finalizeLiveCall(record) + return runErr + } + if !errors.Is(runErr, context.DeadlineExceeded) && time.Now().Before(record.ExpiresAt) { + go s.observeLiveCall(record.CallHash) + return runErr + } + s.finalizeLiveCall(record) + return runErr +} + +func (s *OpenAIGatewayService) runLiveController( + ctx context.Context, + record *LiveCallRecord, + upstream liveFrameConn, + errCh <-chan error, +) error { + refreshTicker := time.NewTicker(liveLeaseRefreshInterval) + defer refreshTicker.Stop() + maxTimer := time.NewTimer(time.Until(record.ExpiresAt)) + defer maxTimer.Stop() + for { + select { + case <-ctx.Done(): + return context.Cause(ctx) + case err := <-errCh: + return err + case <-maxTimer.C: + closeCtx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + _ = upstream.WriteFrame(closeCtx, coderws.MessageText, []byte(`{"type":"session.close"}`)) + cancel() + return context.DeadlineExceeded + case <-refreshTicker.C: + if !s.refreshLiveLease(record) { + return ErrLiveUnavailable + } + } + } +} + +func (s *OpenAIGatewayService) observeLiveCall(callHash string) { + store, err := s.liveStore() + if err != nil { + return + } + owner := uuid.NewString() + claimed, err := store.ClaimLiveController(context.Background(), callHash, LiveControllerObserver, owner) + if err != nil || !claimed { + return + } + for { + record, getErr := store.GetLiveCall(context.Background(), callHash) + if getErr != nil || record.Controller != LiveControllerObserver { + return + } + if !time.Now().Before(record.ExpiresAt) { + s.finalizeLiveCall(record) + return + } + upstream, dialErr := s.dialLiveSideband(context.Background(), record) + if dialErr != nil { + if !s.waitForLiveObserverRetry(record) { + return + } + continue + } + runErr := s.runLiveObserverConnection(record, upstream) + _ = upstream.Close() + if errors.Is(runErr, ErrLiveControllerChanged) { + return + } + if errors.Is(runErr, context.DeadlineExceeded) || errors.Is(runErr, ErrLiveCallNotFound) { + s.finalizeLiveCall(record) + return + } + if !s.waitForLiveObserverRetry(record) { + return + } + } +} + +func (s *OpenAIGatewayService) runLiveObserverConnection(record *LiveCallRecord, upstream liveFrameConn) error { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + frameCh := make(chan []byte, 1) + errCh := make(chan error, 1) + go func() { + for { + messageType, payload, err := upstream.ReadFrame(ctx) + if err != nil { + select { + case errCh <- liveSidebandReadError(err): + case <-ctx.Done(): + } + return + } + if messageType == coderws.MessageText { + select { + case frameCh <- payload: + case <-ctx.Done(): + return + } + } + } + }() + refreshTicker := time.NewTicker(liveLeaseRefreshInterval) + defer refreshTicker.Stop() + controllerTicker := time.NewTicker(liveObserverPollInterval) + defer controllerTicker.Stop() + maxTimer := time.NewTimer(time.Until(record.ExpiresAt)) + defer maxTimer.Stop() + store, _ := s.liveStore() + for { + select { + case payload := <-frameCh: + eventType := strings.TrimSpace(gjson.GetBytes(payload, "type").String()) + if eventType == "session.closed" || eventType == "session.ended" { + return ErrLiveCallNotFound + } + case err := <-errCh: + return err + case <-controllerTicker.C: + controller, err := store.GetLiveController(context.Background(), record.CallHash) + if err != nil { + return err + } + if controller != LiveControllerObserver { + return ErrLiveControllerChanged + } + case <-refreshTicker.C: + if !s.refreshLiveLease(record) { + return ErrLiveUnavailable + } + case <-maxTimer.C: + closeCtx, closeCancel := context.WithTimeout(context.Background(), 2*time.Second) + _ = upstream.WriteFrame(closeCtx, coderws.MessageText, []byte(`{"type":"session.close"}`)) + closeCancel() + return context.DeadlineExceeded + } + } +} + +func (s *OpenAIGatewayService) waitForLiveObserverRetry(record *LiveCallRecord) bool { + timer := time.NewTimer(time.Second) + defer timer.Stop() + <-timer.C + store, err := s.liveStore() + if err != nil { + return false + } + controller, err := store.GetLiveController(context.Background(), record.CallHash) + return err == nil && controller == LiveControllerObserver && time.Now().Before(record.ExpiresAt) +} + +func (s *OpenAIGatewayService) refreshLiveLease(record *LiveCallRecord) bool { + cache, err := s.liveConcurrencyCache() + if err != nil { + return false + } + ctx, cancel := context.WithTimeout(context.Background(), liveRedisOperationTimeout) + defer cancel() + refreshed, err := cache.RefreshLiveLease(ctx, record.AccountID, record.UserID, record.APIKeyID, record.LeaseID) + return err == nil && refreshed +} + +func (s *OpenAIGatewayService) releaseLiveLease(accountID, userID, apiKeyID int64, leaseID string) { + cache, err := s.liveConcurrencyCache() + if err != nil { + return + } + ctx, cancel := context.WithTimeout(context.Background(), liveRedisOperationTimeout) + defer cancel() + _ = cache.ReleaseLiveLease(ctx, accountID, userID, apiKeyID, leaseID) +} + +func (s *OpenAIGatewayService) finalizeLiveCall(record *LiveCallRecord) { + if record == nil { + return + } + store, err := s.liveStore() + if err != nil { + return + } + ctx, cancel := context.WithTimeout(context.Background(), liveRedisOperationTimeout) + first, err := store.MarkLiveCallClosed(ctx, record.CallHash, liveClosedRecordTTL) + cancel() + if err != nil || !first { + return + } + s.releaseLiveLease(record.AccountID, record.UserID, record.APIKeyID, record.LeaseID) + if s.usageLogRepo == nil { + return + } + duration := int(time.Since(record.CreatedAt).Milliseconds()) + if duration < 0 { + duration = 0 + } + inboundEndpoint := record.InboundEndpoint + upstreamEndpoint := "/backend-api/codex/realtime/calls" + userAgent := record.UserAgent + ipAddress := record.IPAddress + billingType := int8(BillingTypeBalance) + if record.SubscriptionID > 0 { + billingType = BillingTypeSubscription + } + _, _ = s.usageLogRepo.Create(context.Background(), &UsageLog{ + UserID: record.UserID, + APIKeyID: record.APIKeyID, + AccountID: record.AccountID, + RequestID: record.CallHash, + Model: record.Model, + RequestedModel: record.Model, + GroupID: liveOptionalID(record.GroupID), + SubscriptionID: liveOptionalID(record.SubscriptionID), + RateMultiplier: 1, + BillingType: billingType, + RequestType: RequestTypeLive, + DurationMs: &duration, + UserAgent: &userAgent, + IPAddress: &ipAddress, + InboundEndpoint: &inboundEndpoint, + UpstreamEndpoint: &upstreamEndpoint, + CreatedAt: record.CreatedAt, + }) +} diff --git a/backend/internal/service/openai_live_lifecycle_test.go b/backend/internal/service/openai_live_lifecycle_test.go new file mode 100644 index 0000000000..e5e071bf4b --- /dev/null +++ b/backend/internal/service/openai_live_lifecycle_test.go @@ -0,0 +1,408 @@ +package service + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + + coderws "github.com/coder/websocket" + "github.com/stretchr/testify/require" +) + +type liveTestFrame struct { + messageType coderws.MessageType + payload []byte + err error +} + +type liveTestFrameConn struct { + reads chan liveTestFrame + writes chan liveTestFrame + closed chan struct{} + closeOnce sync.Once +} + +func newLiveTestFrameConn() *liveTestFrameConn { + return &liveTestFrameConn{ + reads: make(chan liveTestFrame, 8), + writes: make(chan liveTestFrame, 8), + closed: make(chan struct{}), + } +} + +func (c *liveTestFrameConn) ReadFrame(ctx context.Context) (coderws.MessageType, []byte, error) { + select { + case frame := <-c.reads: + return frame.messageType, frame.payload, frame.err + case <-c.closed: + return coderws.MessageText, nil, coderws.CloseError{Code: coderws.StatusNormalClosure} + case <-ctx.Done(): + return coderws.MessageText, nil, context.Cause(ctx) + } +} + +func (c *liveTestFrameConn) WriteFrame(ctx context.Context, messageType coderws.MessageType, payload []byte) error { + frame := liveTestFrame{messageType: messageType, payload: append([]byte(nil), payload...)} + select { + case c.writes <- frame: + return nil + case <-c.closed: + return errors.New("connection closed") + case <-ctx.Done(): + return context.Cause(ctx) + } +} + +func (c *liveTestFrameConn) WriteJSON(ctx context.Context, value any) error { + payload, err := json.Marshal(value) + if err != nil { + return err + } + return c.WriteFrame(ctx, coderws.MessageText, payload) +} + +func (c *liveTestFrameConn) ReadMessage(ctx context.Context) ([]byte, error) { + _, payload, err := c.ReadFrame(ctx) + return payload, err +} + +func (c *liveTestFrameConn) Ping(context.Context) error { return nil } + +func (c *liveTestFrameConn) Close() error { + c.closeOnce.Do(func() { close(c.closed) }) + return nil +} + +type liveTestDialer struct { + conn *liveTestFrameConn + url string + headers http.Header +} + +func (d *liveTestDialer) Dial( + _ context.Context, + wsURL string, + headers http.Header, + _ string, +) (openAIWSClientConn, int, http.Header, error) { + d.url = wsURL + d.headers = headers.Clone() + return d.conn, http.StatusSwitchingProtocols, nil, nil +} + +type liveTestAccountRepo struct { + AccountRepository + account *Account +} + +func (r *liveTestAccountRepo) GetByID(context.Context, int64) (*Account, error) { + return r.account, nil +} + +type liveTestStore struct { + GatewayCache + mu sync.Mutex + record *LiveCallRecord +} + +func (s *liveTestStore) SaveLiveCall(_ context.Context, record *LiveCallRecord, _ time.Duration) error { + s.mu.Lock() + defer s.mu.Unlock() + copy := *record + s.record = © + return nil +} + +func (s *liveTestStore) GetLiveCall(_ context.Context, callHash string) (*LiveCallRecord, error) { + s.mu.Lock() + defer s.mu.Unlock() + if s.record == nil || s.record.CallHash != callHash { + return nil, ErrLiveCallNotFound + } + copy := *s.record + return ©, nil +} + +func (s *liveTestStore) ClaimLiveController(_ context.Context, callHash, controller, owner string) (bool, error) { + s.mu.Lock() + defer s.mu.Unlock() + if s.record == nil || s.record.CallHash != callHash || s.record.Controller == LiveControllerClosed { + return false, nil + } + if controller == LiveControllerObserver && s.record.Controller != LiveControllerPending { + return false, nil + } + if controller == LiveControllerProxy && s.record.Controller != LiveControllerPending && s.record.Controller != LiveControllerObserver { + return false, nil + } + s.record.Controller = controller + s.record.ControllerOwner = owner + return true, nil +} + +func (s *liveTestStore) ReleaseLiveController(_ context.Context, callHash, owner string) (bool, error) { + s.mu.Lock() + defer s.mu.Unlock() + if s.record == nil || s.record.CallHash != callHash || s.record.ControllerOwner != owner { + return false, nil + } + s.record.Controller = LiveControllerPending + s.record.ControllerOwner = "" + return true, nil +} + +func (s *liveTestStore) GetLiveController(_ context.Context, callHash string) (string, error) { + s.mu.Lock() + defer s.mu.Unlock() + if s.record == nil || s.record.CallHash != callHash { + return "", ErrLiveCallNotFound + } + return s.record.Controller, nil +} + +func (s *liveTestStore) MarkLiveCallClosed(_ context.Context, callHash string, _ time.Duration) (bool, error) { + s.mu.Lock() + defer s.mu.Unlock() + if s.record == nil || s.record.CallHash != callHash || s.record.Controller == LiveControllerClosed { + return false, nil + } + s.record.Controller = LiveControllerClosed + s.record.ControllerOwner = "" + return true, nil +} + +type liveTestConcurrencyCache struct { + ConcurrencyCache + mu sync.Mutex + releases int +} + +func (c *liveTestConcurrencyCache) AcquireLiveLease( + context.Context, + int64, + int, + int64, + int, + int64, + string, + bool, +) (bool, error) { + return true, nil +} + +func (c *liveTestConcurrencyCache) RefreshLiveLease( + context.Context, + int64, + int64, + int64, + string, +) (bool, error) { + return true, nil +} + +func (c *liveTestConcurrencyCache) ReleaseLiveLease( + context.Context, + int64, + int64, + int64, + string, +) error { + c.mu.Lock() + c.releases++ + c.mu.Unlock() + return nil +} + +type liveTestUsageRepo struct { + UsageLogRepository + mu sync.Mutex + logs []*UsageLog +} + +func (r *liveTestUsageRepo) Create(_ context.Context, log *UsageLog) (bool, error) { + r.mu.Lock() + defer r.mu.Unlock() + copy := *log + r.logs = append(r.logs, ©) + return true, nil +} + +func TestRunLiveControllerClosesExpiredSession(t *testing.T) { + upstream := newLiveTestFrameConn() + record := &LiveCallRecord{ExpiresAt: time.Now().Add(20 * time.Millisecond)} + service := &OpenAIGatewayService{} + + err := service.runLiveController(context.Background(), record, upstream, make(chan error)) + require.ErrorIs(t, err, context.DeadlineExceeded) + + select { + case frame := <-upstream.writes: + require.Equal(t, coderws.MessageText, frame.messageType) + require.JSONEq(t, `{"type":"session.close"}`, string(frame.payload)) + case <-time.After(time.Second): + t.Fatal("没有向上游发送 session.close") + } +} + +func TestFinalizeLiveCallIsIdempotentAndWritesZeroUsage(t *testing.T) { + record := &LiveCallRecord{ + CallID: "call_secret", + CallHash: hashLiveCallID("call_secret"), + AccountID: 11, + APIKeyID: 22, + UserID: 33, + GroupID: 44, + LeaseID: "lease-1", + Model: "gpt-live-test", + CreatedAt: time.Now().Add(-time.Second), + ExpiresAt: time.Now().Add(time.Hour), + Controller: LiveControllerPending, + InboundEndpoint: "/v1/live", + } + store := &liveTestStore{} + require.NoError(t, store.SaveLiveCall(context.Background(), record, time.Hour)) + concurrencyCache := &liveTestConcurrencyCache{} + usageRepo := &liveTestUsageRepo{} + service := &OpenAIGatewayService{ + cache: store, + concurrencyService: NewConcurrencyService(concurrencyCache), + usageLogRepo: usageRepo, + } + + service.finalizeLiveCall(record) + service.finalizeLiveCall(record) + + concurrencyCache.mu.Lock() + require.Equal(t, 1, concurrencyCache.releases) + concurrencyCache.mu.Unlock() + usageRepo.mu.Lock() + require.Len(t, usageRepo.logs, 1) + log := usageRepo.logs[0] + usageRepo.mu.Unlock() + require.Equal(t, RequestTypeLive, log.RequestType) + require.Equal(t, record.CallHash, log.RequestID) + require.NotEqual(t, record.CallID, log.RequestID) + require.NotNil(t, log.DurationMs) + require.Zero(t, log.InputTokens) + require.Zero(t, log.OutputTokens) + require.Zero(t, log.TotalCost) + require.Zero(t, log.ActualCost) +} + +func TestGetLiveCallForIdentityRejectsMismatchedCaller(t *testing.T) { + groupID := int64(44) + record := &LiveCallRecord{ + CallID: "call_identity", + CallHash: hashLiveCallID("call_identity"), + APIKeyID: 22, + UserID: 33, + GroupID: groupID, + Controller: LiveControllerPending, + } + store := &liveTestStore{} + require.NoError(t, store.SaveLiveCall(context.Background(), record, time.Hour)) + service := &OpenAIGatewayService{cache: store} + + _, err := service.GetLiveCallForIdentity(context.Background(), record.CallID, LiveCallIdentity{ + APIKeyID: 99, + UserID: record.UserID, + GroupID: &groupID, + }) + require.ErrorIs(t, err, ErrLiveIdentityMismatch) + + loaded, err := service.GetLiveCallForIdentity(context.Background(), record.CallID, LiveCallIdentity{ + APIKeyID: record.APIKeyID, + UserID: record.UserID, + GroupID: &groupID, + }) + require.NoError(t, err) + require.Equal(t, record.AccountID, loaded.AccountID) +} + +func TestProxyLiveSidebandForwardsTextAndBinary(t *testing.T) { + account := &Account{ + ID: 11, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Concurrency: 2, + Credentials: map[string]any{ + "access_token": "test-access-token", + "chatgpt_account_id": "acct_test", + }, + } + record := &LiveCallRecord{ + CallID: "call_proxy", + CallHash: hashLiveCallID("call_proxy"), + AccountID: account.ID, + APIKeyID: 22, + UserID: 33, + LeaseID: "lease-1", + CreatedAt: time.Now(), + ExpiresAt: time.Now().Add(time.Minute), + Controller: LiveControllerPending, + } + store := &liveTestStore{} + require.NoError(t, store.SaveLiveCall(context.Background(), record, time.Hour)) + upstream := newLiveTestFrameConn() + dialer := &liveTestDialer{conn: upstream} + service := &OpenAIGatewayService{ + accountRepo: &liveTestAccountRepo{account: account}, + cache: store, + openaiWSPassthroughDialer: dialer, + } + proxyResult := make(chan error, 1) + server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) { + downstream, err := coderws.Accept(writer, request, nil) + if err != nil { + proxyResult <- err + return + } + defer downstream.CloseNow() + proxyResult <- service.ProxyLiveSideband(request.Context(), record, downstream) + })) + defer server.Close() + + client, _, err := coderws.Dial( + context.Background(), + "ws"+strings.TrimPrefix(server.URL, "http"), + nil, + ) + require.NoError(t, err) + defer client.CloseNow() + + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + require.NoError(t, client.Write(ctx, coderws.MessageText, []byte(`{"type":"client.text"}`))) + clientText := <-upstream.writes + require.Equal(t, coderws.MessageText, clientText.messageType) + require.JSONEq(t, `{"type":"client.text"}`, string(clientText.payload)) + + require.NoError(t, client.Write(ctx, coderws.MessageBinary, []byte{1, 2, 3})) + clientBinary := <-upstream.writes + require.Equal(t, coderws.MessageBinary, clientBinary.messageType) + require.Equal(t, []byte{1, 2, 3}, clientBinary.payload) + + upstream.reads <- liveTestFrame{messageType: coderws.MessageText, payload: []byte(`{"type":"server.text"}`)} + messageType, payload, err := client.Read(ctx) + require.NoError(t, err) + require.Equal(t, coderws.MessageText, messageType) + require.JSONEq(t, `{"type":"server.text"}`, string(payload)) + + upstream.reads <- liveTestFrame{messageType: coderws.MessageBinary, payload: []byte{4, 5, 6}} + messageType, payload, err = client.Read(ctx) + require.NoError(t, err) + require.Equal(t, coderws.MessageBinary, messageType) + require.Equal(t, []byte{4, 5, 6}, payload) + + require.Equal(t, "wss://chatgpt.com/backend-api/codex/call_proxy", dialer.url) + require.Equal(t, "Bearer test-access-token", dialer.headers.Get("Authorization")) + require.Equal(t, "acct_test", dialer.headers.Get("Chatgpt-Account-Id")) + upstream.reads <- liveTestFrame{err: coderws.CloseError{Code: coderws.StatusNormalClosure}} + require.ErrorIs(t, <-proxyResult, ErrLiveCallNotFound) +} diff --git a/backend/internal/service/openai_live_test.go b/backend/internal/service/openai_live_test.go new file mode 100644 index 0000000000..9192335880 --- /dev/null +++ b/backend/internal/service/openai_live_test.go @@ -0,0 +1,181 @@ +package service + +import ( + "context" + "encoding/json" + "errors" + "io" + "net/http" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/Wei-Shaw/sub2api/internal/pkg/tlsfingerprint" + coderws "github.com/coder/websocket" + "github.com/stretchr/testify/require" +) + +type liveHTTPUpstreamStub struct { + request *http.Request + body []byte +} + +func (s *liveHTTPUpstreamStub) Do( + request *http.Request, + _ string, + _ int64, + _ int, +) (*http.Response, error) { + s.request = request + body, err := io.ReadAll(request.Body) + if err != nil { + return nil, err + } + s.body = body + return &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Location": {"/backend-api/codex/call_test"}, + }, + Body: io.NopCloser(strings.NewReader("v=0\r\n")), + }, nil +} + +func (s *liveHTTPUpstreamStub) DoWithTLS( + request *http.Request, + proxyURL string, + accountID int64, + accountConcurrency int, + _ *tlsfingerprint.Profile, +) (*http.Response, error) { + return s.Do(request, proxyURL, accountID, accountConcurrency) +} + +func TestLiveCapabilityOnlyAllowsOpenAIOAuth(t *testing.T) { + require.True(t, (&Account{Platform: PlatformOpenAI, Type: AccountTypeOAuth}).SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityLive)) + require.False(t, (&Account{Platform: PlatformOpenAI, Type: AccountTypeAPIKey}).SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityLive)) + require.False(t, (&Account{Platform: PlatformGrok, Type: AccountTypeOAuth}).SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityLive)) + require.False(t, (&Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + openAIAuthModeCredentialKey: OpenAIAuthModePersonalAccessToken, + }, + }).SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityLive)) + require.False(t, (&Account{ + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Credentials: map[string]any{ + openAIAuthModeCredentialKey: OpenAIAuthModeAgentIdentity, + }, + }).SupportsOpenAIEndpointCapability(OpenAIEndpointCapabilityLive)) +} + +func TestValidateLiveCallRequestDoesNotRequireDelegation(t *testing.T) { + request := &LiveCallRequest{ + SDP: "v=0\r\n", + Session: json.RawMessage(`{"model":"gpt-live-test","instructions":"hello"}`), + } + require.NoError(t, ValidateLiveCallRequest(request)) + require.NotContains(t, string(request.Session), "delegation") +} + +func TestCreateUpstreamLiveCallPreservesSession(t *testing.T) { + upstream := &liveHTTPUpstreamStub{} + service := &OpenAIGatewayService{ + cfg: &config.Config{}, + httpUpstream: upstream, + } + account := &Account{ + ID: 7, + Platform: PlatformOpenAI, + Type: AccountTypeOAuth, + Concurrency: 2, + Credentials: map[string]any{ + "access_token": "test-access-token", + "chatgpt_account_id": "acct_test", + }, + } + session := json.RawMessage(`{ + "model":"gpt-live-test", + "delegation":{"type":"client"}, + "custom":{"keep":true} + }`) + + created, err := service.createUpstreamLiveCall(context.Background(), account, &LiveCallRequest{ + SDP: "v=offer\r\n", + Session: session, + }) + require.NoError(t, err) + require.Equal(t, "call_test", created.CallID) + require.Equal(t, []byte("v=0\r\n"), created.SDP) + + var forwarded struct { + SDP string `json:"sdp"` + Session json.RawMessage `json:"session"` + } + require.NoError(t, json.Unmarshal(upstream.body, &forwarded)) + require.Equal(t, "v=offer\r\n", forwarded.SDP) + require.JSONEq(t, string(session), string(forwarded.Session)) + require.Equal(t, "Bearer test-access-token", upstream.request.Header.Get("Authorization")) + require.Equal(t, "acct_test", upstream.request.Header.Get("Chatgpt-Account-Id")) + require.Equal(t, "quicksilver=v2", upstream.request.Header.Get("OpenAI-Alpha")) + require.Empty(t, upstream.request.Header.Get("OpenAI-Beta")) + require.Equal(t, HTTPUpstreamProfileOpenAI, HTTPUpstreamProfileFromContext(upstream.request.Context())) + require.True(t, HTTPUpstreamRedirectsDisabled(upstream.request.Context())) +} + +func TestLiveMaxSessionDurationDefaultsAndOverrides(t *testing.T) { + require.Equal(t, defaultLiveMaxSessionDuration, (&OpenAIGatewayService{}).liveMaxSessionDuration()) + require.Equal( + t, + 90*time.Second, + (&OpenAIGatewayService{cfg: &config.Config{ + Gateway: config.GatewayConfig{ + Live: config.GatewayLiveConfig{MaxSessionDurationSeconds: 90}, + }, + }}).liveMaxSessionDuration(), + ) +} + +func TestLiveSidebandNormalCloseEndsCall(t *testing.T) { + normalClose := coderws.CloseError{Code: coderws.StatusNormalClosure} + require.ErrorIs(t, liveSidebandReadError(normalClose), ErrLiveCallNotFound) + + abnormalClose := coderws.CloseError{Code: coderws.StatusInternalError} + require.Equal(t, abnormalClose, liveSidebandReadError(abnormalClose)) +} + +func TestLiveCreateFailoverUsesExistingOpenAIPolicy(t *testing.T) { + service := &OpenAIGatewayService{} + require.False(t, service.shouldFailoverLiveCreateError(&UpstreamFailoverError{ + StatusCode: http.StatusBadRequest, + ResponseBody: []byte(`{"error":{"message":"invalid session"}}`), + })) + require.True(t, service.shouldFailoverLiveCreateError(&UpstreamFailoverError{ + StatusCode: http.StatusForbidden, + })) + require.True(t, service.shouldFailoverLiveCreateError(&UpstreamFailoverError{ + StatusCode: http.StatusBadGateway, + })) + require.True(t, service.shouldFailoverLiveCreateError(errors.New("transport failed"))) +} + +func TestLiveCallIDFromLocation(t *testing.T) { + callID, err := liveCallIDFromLocation("https://chatgpt.com/backend-api/codex/call_123?intent=quicksilver") + require.NoError(t, err) + require.Equal(t, "call_123", callID) + + callID, err = liveCallIDFromLocation("/backend-api/codex/call_456") + require.NoError(t, err) + require.Equal(t, "call_456", callID) +} + +func TestRequestTypeLive(t *testing.T) { + require.True(t, RequestTypeLive.IsValid()) + require.Equal(t, "live", RequestTypeLive.String()) + parsed, err := ParseUsageRequestType("live") + require.NoError(t, err) + require.Equal(t, RequestTypeLive, parsed) +} diff --git a/backend/internal/service/openai_live_types.go b/backend/internal/service/openai_live_types.go new file mode 100644 index 0000000000..1ada2f362a --- /dev/null +++ b/backend/internal/service/openai_live_types.go @@ -0,0 +1,90 @@ +package service + +import ( + "context" + "encoding/json" + "errors" + "time" +) + +const ( + LiveControllerPending = "pending" + LiveControllerObserver = "observer" + LiveControllerProxy = "proxy" + LiveControllerClosed = "closed" +) + +var ( + ErrLiveUnavailable = errors.New("live is unavailable") + ErrLiveConcurrencyFull = errors.New("live concurrency is full") + ErrLiveCallNotFound = errors.New("live call not found") + ErrLiveIdentityMismatch = errors.New("live call identity mismatch") + ErrLiveControllerChanged = errors.New("live controller changed") +) + +// LiveCallRequest 是两个下游创建协议归一后的请求。Session 不做结构改写。 +type LiveCallRequest struct { + SDP string `json:"sdp"` + Session json.RawMessage `json:"session"` +} + +type LiveCallIdentity struct { + APIKeyID int64 + UserID int64 + GroupID *int64 + SubscriptionID *int64 + UserAgent string + IPAddress string + InboundEndpoint string +} + +type LiveCallRecord struct { + CallID string + CallHash string + AccountID int64 + APIKeyID int64 + UserID int64 + GroupID int64 + SubscriptionID int64 + LeaseID string + Model string + CreatedAt time.Time + ExpiresAt time.Time + Controller string + ControllerOwner string + UserAgent string + IPAddress string + InboundEndpoint string +} + +type LiveCallCreated struct { + SDP []byte + CallID string + Location string + Account *Account +} + +// LiveCallStore 由 GatewayCache 的 Redis 实现可选提供,避免扩大旧缓存接口。 +type LiveCallStore interface { + SaveLiveCall(ctx context.Context, record *LiveCallRecord, ttl time.Duration) error + GetLiveCall(ctx context.Context, callHash string) (*LiveCallRecord, error) + ClaimLiveController(ctx context.Context, callHash, controller, owner string) (bool, error) + ReleaseLiveController(ctx context.Context, callHash, owner string) (bool, error) + GetLiveController(ctx context.Context, callHash string) (string, error) + MarkLiveCallClosed(ctx context.Context, callHash string, ttl time.Duration) (bool, error) +} + +type LiveConcurrencyCache interface { + AcquireLiveLease( + ctx context.Context, + accountID int64, + accountMax int, + userID int64, + userMax int, + apiKeyID int64, + leaseID string, + replacingRegularSlots bool, + ) (bool, error) + RefreshLiveLease(ctx context.Context, accountID, userID, apiKeyID int64, leaseID string) (bool, error) + ReleaseLiveLease(ctx context.Context, accountID, userID, apiKeyID int64, leaseID string) error +} diff --git a/backend/internal/service/usage_log.go b/backend/internal/service/usage_log.go index 4e8df04aaf..f0f586b977 100644 --- a/backend/internal/service/usage_log.go +++ b/backend/internal/service/usage_log.go @@ -19,11 +19,12 @@ const ( RequestTypeStream RequestType = 2 RequestTypeWSV2 RequestType = 3 RequestTypeCyberBlocked RequestType = 4 // cyber_policy 命中(透传但被上游安全策略拒绝) + RequestTypeLive RequestType = 5 ) func (t RequestType) IsValid() bool { switch t { - case RequestTypeUnknown, RequestTypeSync, RequestTypeStream, RequestTypeWSV2, RequestTypeCyberBlocked: + case RequestTypeUnknown, RequestTypeSync, RequestTypeStream, RequestTypeWSV2, RequestTypeCyberBlocked, RequestTypeLive: return true default: return false @@ -47,6 +48,8 @@ func (t RequestType) String() string { return "ws_v2" case RequestTypeCyberBlocked: return "cyber" + case RequestTypeLive: + return "live" default: return "unknown" } @@ -68,8 +71,10 @@ func ParseUsageRequestType(value string) (RequestType, error) { return RequestTypeWSV2, nil case "cyber": return RequestTypeCyberBlocked, nil + case "live": + return RequestTypeLive, nil default: - return RequestTypeUnknown, fmt.Errorf("invalid request_type, allowed values: unknown, sync, stream, ws_v2, cyber") + return RequestTypeUnknown, fmt.Errorf("invalid request_type, allowed values: unknown, sync, stream, ws_v2, cyber, live") } } diff --git a/backend/migrations/186_allow_live_usage_request_type.sql b/backend/migrations/186_allow_live_usage_request_type.sql new file mode 100644 index 0000000000..923315d01f --- /dev/null +++ b/backend/migrations/186_allow_live_usage_request_type.sql @@ -0,0 +1,6 @@ +ALTER TABLE usage_logs + DROP CONSTRAINT IF EXISTS usage_logs_request_type_check; + +ALTER TABLE usage_logs + ADD CONSTRAINT usage_logs_request_type_check + CHECK (request_type >= 0 AND request_type <= 5); diff --git a/backend/migrations/187_add_group_allow_live.sql b/backend/migrations/187_add_group_allow_live.sql new file mode 100644 index 0000000000..660fba2ceb --- /dev/null +++ b/backend/migrations/187_add_group_allow_live.sql @@ -0,0 +1 @@ +ALTER TABLE groups ADD COLUMN allow_live BOOLEAN NOT NULL DEFAULT false; diff --git a/deploy/config.example.yaml b/deploy/config.example.yaml index f66224457a..535a599845 100644 --- a/deploy/config.example.yaml +++ b/deploy/config.example.yaml @@ -293,6 +293,9 @@ gateway: # Use this to avoid compact failures when newer models are not yet supported by the compact endpoint. # 当 compact 端点暂未支持更新模型时,可通过这里降级规避失败。 openai_compact_model: "gpt-5.4" + # ChatGPT Frameless Live 单会话硬上限(秒)。 + live: + max_session_duration_seconds: 3600 # OpenAI Responses WebSocket 配置(默认开启,可按需回滚到 HTTP) openai_ws: # 新版 WS mode 路由(默认关闭)。关闭时忽略账号级 WS mode(包括 http_bridge),保持 legacy ctx_pool 行为。 diff --git a/frontend/src/components/admin/usage/UsageFilters.vue b/frontend/src/components/admin/usage/UsageFilters.vue index 7ac50aa6b2..6a5fad6269 100644 --- a/frontend/src/components/admin/usage/UsageFilters.vue +++ b/frontend/src/components/admin/usage/UsageFilters.vue @@ -263,6 +263,7 @@ const groupOptions = ref([{ value: null, label: t('admin.usage.a const requestTypeOptions = ref([ { value: null, label: t('admin.usage.allTypes') }, { value: 'ws_v2', label: t('usage.ws') }, + { value: 'live', label: t('usage.live') }, { value: 'stream', label: t('usage.stream') }, { value: 'sync', label: t('usage.sync') }, { value: 'cyber', label: t('usage.cyber') } diff --git a/frontend/src/components/admin/usage/UsageTable.vue b/frontend/src/components/admin/usage/UsageTable.vue index ea69fefb2f..032ecc0bd2 100644 --- a/frontend/src/components/admin/usage/UsageTable.vue +++ b/frontend/src/components/admin/usage/UsageTable.vue @@ -583,6 +583,7 @@ const tokenTooltipData = ref(null) const getRequestTypeLabel = (row: AdminUsageLog): string => { const requestType = resolveUsageRequestType(row) if (requestType === 'cyber') return t('usage.cyber') + if (requestType === 'live') return t('usage.live') if (requestType === 'ws_v2') return t('usage.ws') if (requestType === 'stream') return t('usage.stream') if (requestType === 'sync') return t('usage.sync') @@ -592,6 +593,7 @@ const getRequestTypeLabel = (row: AdminUsageLog): string => { const getRequestTypeBadgeClass = (row: AdminUsageLog): string => { const requestType = resolveUsageRequestType(row) if (requestType === 'cyber') return 'bg-red-100 text-red-800 dark:bg-red-900 dark:text-red-200' + if (requestType === 'live') return 'bg-emerald-100 text-emerald-800 dark:bg-emerald-900 dark:text-emerald-200' if (requestType === 'ws_v2') return 'bg-violet-100 text-violet-800 dark:bg-violet-900 dark:text-violet-200' if (requestType === 'stream') return 'bg-blue-100 text-blue-800 dark:bg-blue-900 dark:text-blue-200' if (requestType === 'sync') return 'bg-gray-100 text-gray-800 dark:bg-gray-700 dark:text-gray-200' diff --git a/frontend/src/i18n/locales/en/admin/overview.ts b/frontend/src/i18n/locales/en/admin/overview.ts index 7782f55583..6f1eacd3b4 100644 --- a/frontend/src/i18n/locales/en/admin/overview.ts +++ b/frontend/src/i18n/locales/en/admin/overview.ts @@ -1095,6 +1095,11 @@ export default { targetModelPlaceholder: 'e.g., gpt-5.4', removeExactMapping: 'Remove Exact Mapping' }, + openaiLive: { + title: 'OpenAI Live', + allow: 'Allow Live access', + hint: 'When enabled, API keys in this OpenAI group can create and control Live voice sessions. Disabled by default.' + }, invalidRequestFallback: { title: 'Invalid Request Fallback Group', hint: 'Triggered only when upstream explicitly returns prompt too long. Leave empty to disable fallback.', diff --git a/frontend/src/i18n/locales/en/dashboard.ts b/frontend/src/i18n/locales/en/dashboard.ts index 06d15b9215..870489ea2c 100644 --- a/frontend/src/i18n/locales/en/dashboard.ts +++ b/frontend/src/i18n/locales/en/dashboard.ts @@ -317,6 +317,7 @@ export default { stream: 'Stream', sync: 'Sync', cyber: 'Cyber', + live: 'Live', unknown: 'Unknown', in: 'In', out: 'Out', diff --git a/frontend/src/i18n/locales/zh/admin/overview.ts b/frontend/src/i18n/locales/zh/admin/overview.ts index 39cbe8a264..d76938a050 100644 --- a/frontend/src/i18n/locales/zh/admin/overview.ts +++ b/frontend/src/i18n/locales/zh/admin/overview.ts @@ -1093,6 +1093,11 @@ export default { targetModelPlaceholder: '例如: gpt-5.4', removeExactMapping: '删除精确映射' }, + openaiLive: { + title: 'OpenAI Live', + allow: '允许访问 Live', + hint: '启用后,此 OpenAI 分组的 API Key 可以创建并控制 Live 语音会话。默认关闭。' + }, invalidRequestFallback: { title: '无效请求兜底分组', hint: '仅当上游明确返回 prompt too long 时才会触发,留空表示不兜底', diff --git a/frontend/src/i18n/locales/zh/dashboard.ts b/frontend/src/i18n/locales/zh/dashboard.ts index eaf0004086..7d129df336 100644 --- a/frontend/src/i18n/locales/zh/dashboard.ts +++ b/frontend/src/i18n/locales/zh/dashboard.ts @@ -322,6 +322,7 @@ export default { stream: '流式', sync: '同步', cyber: '安全策略', + live: 'Live', unknown: '未知', in: '输入', out: '输出', diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts index a96180c1de..6198d8f832 100644 --- a/frontend/src/types/index.ts +++ b/frontend/src/types/index.ts @@ -551,6 +551,8 @@ export interface Group { fallback_group_id_on_invalid_request: number | null // OpenAI Messages 调度开关(用户侧需要此字段判断是否展示 Claude Code 教程) allow_messages_dispatch?: boolean + // OpenAI Live 接口开关 + allow_live: boolean default_mapped_model?: string messages_dispatch_model_config?: OpenAIMessagesDispatchModelConfig require_oauth_only: boolean @@ -742,6 +744,7 @@ export interface CreateGroupRequest { supported_model_scopes?: string[] models_list_config?: ModelsListConfig allow_messages_dispatch?: boolean + allow_live?: boolean default_mapped_model?: string messages_dispatch_model_config?: OpenAIMessagesDispatchModelConfig model_routing?: Record | null @@ -792,6 +795,7 @@ export interface UpdateGroupRequest { supported_model_scopes?: string[] models_list_config?: ModelsListConfig allow_messages_dispatch?: boolean + allow_live?: boolean default_mapped_model?: string messages_dispatch_model_config?: OpenAIMessagesDispatchModelConfig model_routing?: Record | null @@ -1511,7 +1515,7 @@ export interface CodexSessionImportResult { // ==================== Usage & Redeem Types ==================== export type RedeemCodeType = 'balance' | 'concurrency' | 'subscription' | 'invitation' -export type UsageRequestType = 'unknown' | 'sync' | 'stream' | 'ws_v2' | 'cyber' +export type UsageRequestType = 'unknown' | 'sync' | 'stream' | 'ws_v2' | 'cyber' | 'live' export type ImageSizeSource = 'output' | 'input' | 'default' | 'legacy' export type ImageSizeBreakdown = Record diff --git a/frontend/src/utils/errorBadges.ts b/frontend/src/utils/errorBadges.ts index bdb4657547..3354f096f3 100644 --- a/frontend/src/utils/errorBadges.ts +++ b/frontend/src/utils/errorBadges.ts @@ -15,9 +15,10 @@ export function statusCodeBadgeClass(code: number): string { return 'bg-gray-100 text-gray-800 dark:bg-dark-700 dark:text-gray-200' } -/** 请求类型徽章配色(cyber 红、ws 紫、stream 蓝、sync 灰、未知琥珀) */ +/** 请求类型徽章配色(cyber 红、live 绿、ws 紫、stream 蓝、sync 灰、未知琥珀) */ export function requestTypeBadgeClass(kind: UsageRequestKind): string { if (kind === 'cyber') return 'bg-red-100 text-red-800 dark:bg-red-900 dark:text-red-200' + if (kind === 'live') return 'bg-emerald-100 text-emerald-800 dark:bg-emerald-900 dark:text-emerald-200' if (kind === 'ws_v2') return 'bg-violet-100 text-violet-800 dark:bg-violet-900 dark:text-violet-200' if (kind === 'stream') return 'bg-blue-100 text-blue-800 dark:bg-blue-900 dark:text-blue-200' if (kind === 'sync') return 'bg-gray-100 text-gray-800 dark:bg-dark-700 dark:text-gray-200' @@ -27,6 +28,7 @@ export function requestTypeBadgeClass(kind: UsageRequestKind): string { /** 请求类型 i18n 键(展示方自行 t()) */ export function requestTypeLabelKey(kind: UsageRequestKind): string { if (kind === 'cyber') return 'usage.cyber' + if (kind === 'live') return 'usage.live' if (kind === 'ws_v2') return 'usage.ws' if (kind === 'stream') return 'usage.stream' if (kind === 'sync') return 'usage.sync' @@ -43,6 +45,7 @@ export function numericRequestTypeKind( ): UsageRequestKind | null { const rt = requestType ?? (stream == null ? 0 : stream ? 2 : 1) if (rt === 3) return 'ws_v2' + if (rt === 5) return 'live' if (rt === 2) return 'stream' if (rt === 1) return 'sync' return null diff --git a/frontend/src/utils/usageRequestType.ts b/frontend/src/utils/usageRequestType.ts index 14de5f043d..13aeaf68dc 100644 --- a/frontend/src/utils/usageRequestType.ts +++ b/frontend/src/utils/usageRequestType.ts @@ -6,7 +6,7 @@ export interface UsageRequestTypeLike { openai_ws_mode?: boolean | null } -const VALID_REQUEST_TYPES = new Set(['unknown', 'sync', 'stream', 'ws_v2', 'cyber']) +const VALID_REQUEST_TYPES = new Set(['unknown', 'sync', 'stream', 'ws_v2', 'cyber', 'live']) export const isUsageRequestType = (value: unknown): value is UsageRequestType => { return typeof value === 'string' && VALID_REQUEST_TYPES.has(value as UsageRequestType) @@ -24,7 +24,7 @@ export const resolveUsageRequestType = (value: UsageRequestTypeLike): UsageReque export const requestTypeToLegacyStream = (requestType?: UsageRequestType | null): boolean | null | undefined => { // cyber 与 stream 正交(cyber 可发生在 stream 或非 stream 请求),不映射到 legacy stream 维度。 - if (!requestType || requestType === 'unknown' || requestType === 'cyber') { + if (!requestType || requestType === 'unknown' || requestType === 'cyber' || requestType === 'live') { return null } if (requestType === 'sync') { diff --git a/frontend/src/views/admin/GroupsView.vue b/frontend/src/views/admin/GroupsView.vue index 6ebe04285f..d56e5d4725 100644 --- a/frontend/src/views/admin/GroupsView.vue +++ b/frontend/src/views/admin/GroupsView.vue @@ -1394,6 +1394,39 @@ + +
+

+ {{ t("admin.groups.openaiLive.title") }} +

+
+ + +
+

+ {{ t("admin.groups.openaiLive.hint") }} +

+
+
+ +
+

+ {{ t("admin.groups.openaiLive.title") }} +

+
+ + +
+

+ {{ t("admin.groups.openaiLive.hint") }} +

+
+
{ createForm.fallback_group_id = null; createForm.fallback_group_id_on_invalid_request = null; resetMessagesDispatchFormState(createForm); + createForm.allow_live = false; createForm.require_oauth_only = false; createForm.require_privacy_set = false; createForm.supported_model_scopes = ["claude", "gemini_text", "gemini_image"]; @@ -5474,6 +5543,7 @@ const handleEdit = async (group: AdminGroup) => { editForm.allow_messages_dispatch = group.allow_messages_dispatch || messagesDispatchFormState.allow_messages_dispatch; + editForm.allow_live = group.allow_live ?? false; editForm.opus_mapped_model = messagesDispatchFormState.opus_mapped_model; editForm.sonnet_mapped_model = messagesDispatchFormState.sonnet_mapped_model; editForm.haiku_mapped_model = messagesDispatchFormState.haiku_mapped_model; @@ -5530,6 +5600,7 @@ const closeEditModal = () => { editForm.video_price_1080p = null; editForm.web_search_price_per_call = null; resetMessagesDispatchFormState(editForm); + editForm.allow_live = false; resetModelsListState(editModelsListState); }; @@ -5934,6 +6005,7 @@ watch( } if (newVal !== "openai") { resetMessagesDispatchFormState(createForm); + createForm.allow_live = false; } createForm.max_reasoning_effort = normalizeReasoningEffortForPlatform( newVal, @@ -5976,6 +6048,7 @@ watch( } if (newVal !== "openai") { resetMessagesDispatchFormState(editForm); + editForm.allow_live = false; } editForm.max_reasoning_effort = normalizeReasoningEffortForPlatform( newVal, @@ -6020,6 +6093,7 @@ watch( } if (newVal !== 'openai') { editForm.allow_messages_dispatch = false + editForm.allow_live = false editForm.default_mapped_model = '' } } diff --git a/frontend/src/views/admin/UsageView.vue b/frontend/src/views/admin/UsageView.vue index 5df00c2898..8930cf635e 100644 --- a/frontend/src/views/admin/UsageView.vue +++ b/frontend/src/views/admin/UsageView.vue @@ -538,6 +538,7 @@ const openCleanupDialog = () => { cleanupDialogVisible.value = true } const getRequestTypeLabel = (log: AdminUsageLog): string => { const requestType = resolveUsageRequestType(log) if (requestType === 'cyber') return t('usage.cyber') + if (requestType === 'live') return t('usage.live') if (requestType === 'ws_v2') return t('usage.ws') if (requestType === 'stream') return t('usage.stream') if (requestType === 'sync') return t('usage.sync') diff --git a/frontend/src/views/user/UsageView.vue b/frontend/src/views/user/UsageView.vue index 8fad943adf..aa788a2b04 100644 --- a/frontend/src/views/user/UsageView.vue +++ b/frontend/src/views/user/UsageView.vue @@ -376,6 +376,7 @@ const granularityOptions = computed(() => [ const requestTypeOptions = computed(() => [ { value: null, label: t('admin.usage.allTypes') }, { value: 'ws_v2', label: t('usage.ws') }, + { value: 'live', label: t('usage.live') }, { value: 'stream', label: t('usage.stream') }, { value: 'sync', label: t('usage.sync') }, ]) @@ -595,6 +596,7 @@ const handleIpGeoBatchFailed = () => { const getRequestTypeExportText = (log: UsageLog): string => { const requestType = resolveUsageRequestType(log) if (requestType === 'cyber') return 'Cyber' + if (requestType === 'live') return 'Live' if (requestType === 'ws_v2') return 'WS' if (requestType === 'stream') return 'Stream' if (requestType === 'sync') return 'Sync'