mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:28:39 +08:00
feat(openai): add Live gateway support
This commit is contained in:
+12
-1
@@ -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(", ")
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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 类型账号关联到此分组"),
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"`
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"},
|
||||
|
||||
@@ -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 上游
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 = ©
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *liveTestStore) GetLiveCall(_ context.Context, callHash string) (*LiveCallRecord, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.record == nil || s.record.CallHash != callHash {
|
||||
return nil, ErrLiveCallNotFound
|
||||
}
|
||||
copy := *s.record
|
||||
return ©, nil
|
||||
}
|
||||
|
||||
func (s *liveTestStore) ClaimLiveController(_ context.Context, callHash, controller, owner string) (bool, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.record == nil || s.record.CallHash != callHash || s.record.Controller == LiveControllerClosed {
|
||||
return false, nil
|
||||
}
|
||||
if controller == LiveControllerObserver && s.record.Controller != LiveControllerPending {
|
||||
return false, nil
|
||||
}
|
||||
if controller == LiveControllerProxy && s.record.Controller != LiveControllerPending && s.record.Controller != LiveControllerObserver {
|
||||
return false, nil
|
||||
}
|
||||
s.record.Controller = controller
|
||||
s.record.ControllerOwner = owner
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (s *liveTestStore) ReleaseLiveController(_ context.Context, callHash, owner string) (bool, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.record == nil || s.record.CallHash != callHash || s.record.ControllerOwner != owner {
|
||||
return false, nil
|
||||
}
|
||||
s.record.Controller = LiveControllerPending
|
||||
s.record.ControllerOwner = ""
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (s *liveTestStore) GetLiveController(_ context.Context, callHash string) (string, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.record == nil || s.record.CallHash != callHash {
|
||||
return "", ErrLiveCallNotFound
|
||||
}
|
||||
return s.record.Controller, nil
|
||||
}
|
||||
|
||||
func (s *liveTestStore) MarkLiveCallClosed(_ context.Context, callHash string, _ time.Duration) (bool, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.record == nil || s.record.CallHash != callHash || s.record.Controller == LiveControllerClosed {
|
||||
return false, nil
|
||||
}
|
||||
s.record.Controller = LiveControllerClosed
|
||||
s.record.ControllerOwner = ""
|
||||
return true, nil
|
||||
}
|
||||
|
||||
type liveTestConcurrencyCache struct {
|
||||
ConcurrencyCache
|
||||
mu sync.Mutex
|
||||
releases int
|
||||
}
|
||||
|
||||
func (c *liveTestConcurrencyCache) AcquireLiveLease(
|
||||
context.Context,
|
||||
int64,
|
||||
int,
|
||||
int64,
|
||||
int,
|
||||
int64,
|
||||
string,
|
||||
bool,
|
||||
) (bool, error) {
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (c *liveTestConcurrencyCache) RefreshLiveLease(
|
||||
context.Context,
|
||||
int64,
|
||||
int64,
|
||||
int64,
|
||||
string,
|
||||
) (bool, error) {
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func (c *liveTestConcurrencyCache) ReleaseLiveLease(
|
||||
context.Context,
|
||||
int64,
|
||||
int64,
|
||||
int64,
|
||||
string,
|
||||
) error {
|
||||
c.mu.Lock()
|
||||
c.releases++
|
||||
c.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
|
||||
type liveTestUsageRepo struct {
|
||||
UsageLogRepository
|
||||
mu sync.Mutex
|
||||
logs []*UsageLog
|
||||
}
|
||||
|
||||
func (r *liveTestUsageRepo) Create(_ context.Context, log *UsageLog) (bool, error) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
copy := *log
|
||||
r.logs = append(r.logs, ©)
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func TestRunLiveControllerClosesExpiredSession(t *testing.T) {
|
||||
upstream := newLiveTestFrameConn()
|
||||
record := &LiveCallRecord{ExpiresAt: time.Now().Add(20 * time.Millisecond)}
|
||||
service := &OpenAIGatewayService{}
|
||||
|
||||
err := service.runLiveController(context.Background(), record, upstream, make(chan error))
|
||||
require.ErrorIs(t, err, context.DeadlineExceeded)
|
||||
|
||||
select {
|
||||
case frame := <-upstream.writes:
|
||||
require.Equal(t, coderws.MessageText, frame.messageType)
|
||||
require.JSONEq(t, `{"type":"session.close"}`, string(frame.payload))
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("没有向上游发送 session.close")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFinalizeLiveCallIsIdempotentAndWritesZeroUsage(t *testing.T) {
|
||||
record := &LiveCallRecord{
|
||||
CallID: "call_secret",
|
||||
CallHash: hashLiveCallID("call_secret"),
|
||||
AccountID: 11,
|
||||
APIKeyID: 22,
|
||||
UserID: 33,
|
||||
GroupID: 44,
|
||||
LeaseID: "lease-1",
|
||||
Model: "gpt-live-test",
|
||||
CreatedAt: time.Now().Add(-time.Second),
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
Controller: LiveControllerPending,
|
||||
InboundEndpoint: "/v1/live",
|
||||
}
|
||||
store := &liveTestStore{}
|
||||
require.NoError(t, store.SaveLiveCall(context.Background(), record, time.Hour))
|
||||
concurrencyCache := &liveTestConcurrencyCache{}
|
||||
usageRepo := &liveTestUsageRepo{}
|
||||
service := &OpenAIGatewayService{
|
||||
cache: store,
|
||||
concurrencyService: NewConcurrencyService(concurrencyCache),
|
||||
usageLogRepo: usageRepo,
|
||||
}
|
||||
|
||||
service.finalizeLiveCall(record)
|
||||
service.finalizeLiveCall(record)
|
||||
|
||||
concurrencyCache.mu.Lock()
|
||||
require.Equal(t, 1, concurrencyCache.releases)
|
||||
concurrencyCache.mu.Unlock()
|
||||
usageRepo.mu.Lock()
|
||||
require.Len(t, usageRepo.logs, 1)
|
||||
log := usageRepo.logs[0]
|
||||
usageRepo.mu.Unlock()
|
||||
require.Equal(t, RequestTypeLive, log.RequestType)
|
||||
require.Equal(t, record.CallHash, log.RequestID)
|
||||
require.NotEqual(t, record.CallID, log.RequestID)
|
||||
require.NotNil(t, log.DurationMs)
|
||||
require.Zero(t, log.InputTokens)
|
||||
require.Zero(t, log.OutputTokens)
|
||||
require.Zero(t, log.TotalCost)
|
||||
require.Zero(t, log.ActualCost)
|
||||
}
|
||||
|
||||
func TestGetLiveCallForIdentityRejectsMismatchedCaller(t *testing.T) {
|
||||
groupID := int64(44)
|
||||
record := &LiveCallRecord{
|
||||
CallID: "call_identity",
|
||||
CallHash: hashLiveCallID("call_identity"),
|
||||
APIKeyID: 22,
|
||||
UserID: 33,
|
||||
GroupID: groupID,
|
||||
Controller: LiveControllerPending,
|
||||
}
|
||||
store := &liveTestStore{}
|
||||
require.NoError(t, store.SaveLiveCall(context.Background(), record, time.Hour))
|
||||
service := &OpenAIGatewayService{cache: store}
|
||||
|
||||
_, err := service.GetLiveCallForIdentity(context.Background(), record.CallID, LiveCallIdentity{
|
||||
APIKeyID: 99,
|
||||
UserID: record.UserID,
|
||||
GroupID: &groupID,
|
||||
})
|
||||
require.ErrorIs(t, err, ErrLiveIdentityMismatch)
|
||||
|
||||
loaded, err := service.GetLiveCallForIdentity(context.Background(), record.CallID, LiveCallIdentity{
|
||||
APIKeyID: record.APIKeyID,
|
||||
UserID: record.UserID,
|
||||
GroupID: &groupID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, record.AccountID, loaded.AccountID)
|
||||
}
|
||||
|
||||
func TestProxyLiveSidebandForwardsTextAndBinary(t *testing.T) {
|
||||
account := &Account{
|
||||
ID: 11,
|
||||
Platform: PlatformOpenAI,
|
||||
Type: AccountTypeOAuth,
|
||||
Concurrency: 2,
|
||||
Credentials: map[string]any{
|
||||
"access_token": "test-access-token",
|
||||
"chatgpt_account_id": "acct_test",
|
||||
},
|
||||
}
|
||||
record := &LiveCallRecord{
|
||||
CallID: "call_proxy",
|
||||
CallHash: hashLiveCallID("call_proxy"),
|
||||
AccountID: account.ID,
|
||||
APIKeyID: 22,
|
||||
UserID: 33,
|
||||
LeaseID: "lease-1",
|
||||
CreatedAt: time.Now(),
|
||||
ExpiresAt: time.Now().Add(time.Minute),
|
||||
Controller: LiveControllerPending,
|
||||
}
|
||||
store := &liveTestStore{}
|
||||
require.NoError(t, store.SaveLiveCall(context.Background(), record, time.Hour))
|
||||
upstream := newLiveTestFrameConn()
|
||||
dialer := &liveTestDialer{conn: upstream}
|
||||
service := &OpenAIGatewayService{
|
||||
accountRepo: &liveTestAccountRepo{account: account},
|
||||
cache: store,
|
||||
openaiWSPassthroughDialer: dialer,
|
||||
}
|
||||
proxyResult := make(chan error, 1)
|
||||
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
|
||||
downstream, err := coderws.Accept(writer, request, nil)
|
||||
if err != nil {
|
||||
proxyResult <- err
|
||||
return
|
||||
}
|
||||
defer downstream.CloseNow()
|
||||
proxyResult <- service.ProxyLiveSideband(request.Context(), record, downstream)
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
client, _, err := coderws.Dial(
|
||||
context.Background(),
|
||||
"ws"+strings.TrimPrefix(server.URL, "http"),
|
||||
nil,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
defer client.CloseNow()
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
defer cancel()
|
||||
require.NoError(t, client.Write(ctx, coderws.MessageText, []byte(`{"type":"client.text"}`)))
|
||||
clientText := <-upstream.writes
|
||||
require.Equal(t, coderws.MessageText, clientText.messageType)
|
||||
require.JSONEq(t, `{"type":"client.text"}`, string(clientText.payload))
|
||||
|
||||
require.NoError(t, client.Write(ctx, coderws.MessageBinary, []byte{1, 2, 3}))
|
||||
clientBinary := <-upstream.writes
|
||||
require.Equal(t, coderws.MessageBinary, clientBinary.messageType)
|
||||
require.Equal(t, []byte{1, 2, 3}, clientBinary.payload)
|
||||
|
||||
upstream.reads <- liveTestFrame{messageType: coderws.MessageText, payload: []byte(`{"type":"server.text"}`)}
|
||||
messageType, payload, err := client.Read(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, coderws.MessageText, messageType)
|
||||
require.JSONEq(t, `{"type":"server.text"}`, string(payload))
|
||||
|
||||
upstream.reads <- liveTestFrame{messageType: coderws.MessageBinary, payload: []byte{4, 5, 6}}
|
||||
messageType, payload, err = client.Read(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, coderws.MessageBinary, messageType)
|
||||
require.Equal(t, []byte{4, 5, 6}, payload)
|
||||
|
||||
require.Equal(t, "wss://chatgpt.com/backend-api/codex/call_proxy", dialer.url)
|
||||
require.Equal(t, "Bearer test-access-token", dialer.headers.Get("Authorization"))
|
||||
require.Equal(t, "acct_test", dialer.headers.Get("Chatgpt-Account-Id"))
|
||||
upstream.reads <- liveTestFrame{err: coderws.CloseError{Code: coderws.StatusNormalClosure}}
|
||||
require.ErrorIs(t, <-proxyResult, ErrLiveCallNotFound)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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;
|
||||
@@ -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: '输出',
|
||||
|
||||
@@ -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>
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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') {
|
||||
|
||||
@@ -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 = ''
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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')
|
||||
|
||||
@@ -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'
|
||||
|
||||
Reference in New Issue
Block a user