feat(openai): add Live gateway support

This commit is contained in:
song
2026-07-25 12:50:46 +08:00
parent 37ed639d1e
commit e6eb23eaac
55 changed files with 2607 additions and 30 deletions
+12 -1
View File
@@ -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(", ")
+10
View File
@@ -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()
+15
View File
@@ -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))
+65
View File
@@ -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) {
+34
View File
@@ -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)
}
+1
View File
@@ -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: ""},
+55 -1
View File
@@ -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
+12 -8
View File
@@ -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()
+3
View File
@@ -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 类型账号关联到此分组"),
+10
View File
@@ -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:
@@ -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,
+1
View File
@@ -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,
+2
View File
@@ -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"`
+231
View File
@@ -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
}
@@ -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
}
@@ -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,
+152 -13
View File
@@ -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
@@ -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"
@@ -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)
}
@@ -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
}
@@ -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)
}
@@ -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).
@@ -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)
@@ -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,
@@ -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)
@@ -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"},
+6
View File
@@ -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 上游
+10
View File
@@ -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 {
@@ -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,
@@ -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",
@@ -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
@@ -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)
}
@@ -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"`
@@ -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,
+1
View File
@@ -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
+747
View File
@@ -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,
})
}
@@ -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 = &copy
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 &copy, 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, &copy)
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)
}
@@ -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)
}
@@ -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
}
+7 -2
View File
@@ -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")
}
}
@@ -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);
@@ -0,0 +1 @@
ALTER TABLE groups ADD COLUMN allow_live BOOLEAN NOT NULL DEFAULT false;
+3
View File
@@ -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 行为。
@@ -263,6 +263,7 @@ const groupOptions = ref<SelectOption[]>([{ value: null, label: t('admin.usage.a
const requestTypeOptions = ref<SelectOption[]>([
{ 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') }
@@ -583,6 +583,7 @@ const tokenTooltipData = ref<AdminUsageLog | null>(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'
@@ -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.',
@@ -317,6 +317,7 @@ export default {
stream: 'Stream',
sync: 'Sync',
cyber: 'Cyber',
live: 'Live',
unknown: 'Unknown',
in: 'In',
out: 'Out',
@@ -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 时才会触发,留空表示不兜底',
@@ -322,6 +322,7 @@ export default {
stream: '流式',
sync: '同步',
cyber: '安全策略',
live: 'Live',
unknown: '未知',
in: '输入',
out: '输出',
+5 -1
View File
@@ -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<string, number[]> | 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<string, number[]> | 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<string, number>
+4 -1
View File
@@ -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
+2 -2
View File
@@ -6,7 +6,7 @@ export interface UsageRequestTypeLike {
openai_ws_mode?: boolean | null
}
const VALID_REQUEST_TYPES = new Set<UsageRequestType>(['unknown', 'sync', 'stream', 'ws_v2', 'cyber'])
const VALID_REQUEST_TYPES = new Set<UsageRequestType>(['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') {
+74
View File
@@ -1394,6 +1394,39 @@
</div>
</div>
<!-- OpenAI Live 开关(仅 openai 平台) -->
<div
v-if="createForm.platform === 'openai'"
class="border-t border-gray-200 dark:border-dark-400 pt-4 mt-4"
>
<h4 class="text-sm font-medium text-gray-700 dark:text-gray-300 mb-3">
{{ t("admin.groups.openaiLive.title") }}
</h4>
<div class="flex items-center justify-between">
<label class="text-sm text-gray-600 dark:text-gray-400">{{
t("admin.groups.openaiLive.allow")
}}</label>
<button
type="button"
@click="createForm.allow_live = !createForm.allow_live"
class="relative inline-flex h-6 w-12 flex-shrink-0 cursor-pointer rounded-full border-2 border-transparent transition-colors duration-200 ease-in-out focus:outline-none"
:class="
createForm.allow_live
? 'bg-primary-500'
: 'bg-gray-300 dark:bg-dark-600'
"
>
<span
class="pointer-events-none inline-block h-5 w-5 transform rounded-full bg-white shadow ring-0 transition duration-200 ease-in-out"
:class="createForm.allow_live ? 'translate-x-6' : 'translate-x-1'"
/>
</button>
</div>
<p class="text-xs text-gray-500 dark:text-gray-400 mt-1">
{{ t("admin.groups.openaiLive.hint") }}
</p>
</div>
<!-- OpenAI Messages 调度配置(仅 openai 平台) -->
<div
v-if="createForm.platform === 'openai'"
@@ -2912,6 +2945,39 @@
</div>
</div>
<!-- OpenAI Live 开关(仅 openai 平台) -->
<div
v-if="editForm.platform === 'openai'"
class="border-t border-gray-200 dark:border-dark-400 pt-4 mt-4"
>
<h4 class="text-sm font-medium text-gray-700 dark:text-gray-300 mb-3">
{{ t("admin.groups.openaiLive.title") }}
</h4>
<div class="flex items-center justify-between">
<label class="text-sm text-gray-600 dark:text-gray-400">{{
t("admin.groups.openaiLive.allow")
}}</label>
<button
type="button"
@click="editForm.allow_live = !editForm.allow_live"
class="relative inline-flex h-6 w-12 flex-shrink-0 cursor-pointer rounded-full border-2 border-transparent transition-colors duration-200 ease-in-out focus:outline-none"
:class="
editForm.allow_live
? 'bg-primary-500'
: 'bg-gray-300 dark:bg-dark-600'
"
>
<span
class="pointer-events-none inline-block h-5 w-5 transform rounded-full bg-white shadow ring-0 transition duration-200 ease-in-out"
:class="editForm.allow_live ? 'translate-x-6' : 'translate-x-1'"
/>
</button>
</div>
<p class="text-xs text-gray-500 dark:text-gray-400 mt-1">
{{ t("admin.groups.openaiLive.hint") }}
</p>
</div>
<!-- OpenAI Messages 调度配置(仅 openai 平台) -->
<div
v-if="editForm.platform === 'openai'"
@@ -4535,6 +4601,7 @@ const createForm = reactive({
fallback_group_id_on_invalid_request: null as number | null,
// OpenAI Messages 调度配置(仅 openai 平台使用)
allow_messages_dispatch: false,
allow_live: false,
opus_mapped_model: createMessagesDispatchDefaults.opus_mapped_model,
sonnet_mapped_model: createMessagesDispatchDefaults.sonnet_mapped_model,
haiku_mapped_model: createMessagesDispatchDefaults.haiku_mapped_model,
@@ -4884,6 +4951,7 @@ const editForm = reactive({
fallback_group_id_on_invalid_request: null as number | null,
// OpenAI Messages 调度配置(仅 openai 平台使用)
allow_messages_dispatch: false,
allow_live: false,
default_mapped_model: '',
opus_mapped_model: editMessagesDispatchDefaults.opus_mapped_model,
sonnet_mapped_model: editMessagesDispatchDefaults.sonnet_mapped_model,
@@ -5287,6 +5355,7 @@ const closeCreateModal = () => {
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 = ''
}
}
+1
View File
@@ -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')
+2
View File
@@ -376,6 +376,7 @@ const granularityOptions = computed<SelectOption[]>(() => [
const requestTypeOptions = computed<SelectOption[]>(() => [
{ 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'