mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 16:08:02 +08:00
Merge pull request #4843 from slovx2/feature/openai-live-gateway
feat(openai): 支持 ChatGPT Live / Frameless 网关
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 {
|
||||
@@ -2188,6 +2195,7 @@ func setDefaults() {
|
||||
viper.SetDefault("gateway.codex_image_generation_bridge_enabled", false)
|
||||
viper.SetDefault("gateway.openai_passthrough_allow_timeout_headers", false)
|
||||
viper.SetDefault("gateway.openai_compact_model", "gpt-5.4")
|
||||
viper.SetDefault("gateway.live.max_session_duration_seconds", 3600)
|
||||
// OpenAI Responses WebSocket(默认开启;可通过 force_http 紧急回滚)
|
||||
viper.SetDefault("gateway.openai_ws.enabled", true)
|
||||
viper.SetDefault("gateway.openai_ws.mode_router_v2_enabled", false)
|
||||
@@ -3039,6 +3047,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:
|
||||
|
||||
@@ -13,6 +13,7 @@ import (
|
||||
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/response"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/timezone"
|
||||
"github.com/Wei-Shaw/sub2api/internal/platform/liveattestation"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -25,6 +26,16 @@ type GroupHandler struct {
|
||||
groupCapacityService *service.GroupCapacityService
|
||||
}
|
||||
|
||||
// GetLiveCapability 返回当前服务端是否具备生成 Live attestation 的运行环境。
|
||||
func (h *GroupHandler) GetLiveCapability(c *gin.Context) {
|
||||
err := liveattestation.NewProvider().Check(c.Request.Context())
|
||||
result := gin.H{"supported": err == nil}
|
||||
if err != nil {
|
||||
result["reason"] = err.Error()
|
||||
}
|
||||
response.Success(c, result)
|
||||
}
|
||||
|
||||
type optionalLimitField struct {
|
||||
set bool
|
||||
value *float64
|
||||
@@ -125,6 +136,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 +195,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 +512,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 +631,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,236 @@
|
||||
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 attestationErr *service.LiveAttestationUnavailableError
|
||||
if errors.As(err, &attestationErr) {
|
||||
h.errorResponse(c, http.StatusServiceUnavailable, "api_error", attestationErr.Error())
|
||||
return
|
||||
}
|
||||
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 func() { _ = 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,118 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"mime/multipart"
|
||||
"net/http"
|
||||
"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 TestLiveAttestationErrorIsExplicit(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
recorder := httptest.NewRecorder()
|
||||
context, _ := gin.CreateTestContext(recorder)
|
||||
|
||||
(&OpenAIGatewayHandler{}).writeLiveCreateError(context, &service.LiveAttestationUnavailableError{
|
||||
Reason: "Live attestation is only supported when Sub2API runs on macOS",
|
||||
})
|
||||
|
||||
require.Equal(t, http.StatusServiceUnavailable, recorder.Code)
|
||||
require.Contains(t, recorder.Body.String(), "Sub2API runs on macOS")
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
package liveattestation
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrUnsupportedPlatform = errors.New("live attestation is only supported when Sub2API runs on macOS; Windows support is not implemented yet")
|
||||
ErrChatGPTAppMissing = errors.New("live attestation requires the official ChatGPT app on the Sub2API server")
|
||||
)
|
||||
|
||||
// Provider 在发起 Live 请求前生成 ChatGPT DeviceCheck attestation。
|
||||
type Provider interface {
|
||||
Check(ctx context.Context) error
|
||||
Generate(ctx context.Context) (string, error)
|
||||
}
|
||||
@@ -0,0 +1,294 @@
|
||||
//go:build darwin
|
||||
|
||||
package liveattestation
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/google/uuid"
|
||||
)
|
||||
|
||||
const (
|
||||
chatGPTApplicationPath = "/Applications/ChatGPT.app"
|
||||
attestationTimeout = 5 * time.Second
|
||||
)
|
||||
|
||||
type darwinProvider struct {
|
||||
appSessionID string
|
||||
appPaths []string
|
||||
}
|
||||
|
||||
type deviceSignals struct {
|
||||
SchemaVersion int `json:"schemaVersion"`
|
||||
PreferredLanguages []string `json:"preferredLanguages"`
|
||||
Locale string `json:"locale"`
|
||||
Timezone string `json:"timezone"`
|
||||
ScreenSizeSum int `json:"screenSizeSum"`
|
||||
ScreenScale float64 `json:"screenScale"`
|
||||
AppSessionID string `json:"appSessionId"`
|
||||
}
|
||||
|
||||
type macOSSignals struct {
|
||||
Locale string `json:"locale"`
|
||||
Languages []string `json:"languages"`
|
||||
Timezone string `json:"timezone"`
|
||||
Width float64 `json:"width"`
|
||||
Height float64 `json:"height"`
|
||||
Scale float64 `json:"scale"`
|
||||
}
|
||||
|
||||
func NewProvider() Provider {
|
||||
paths := []string{chatGPTApplicationPath}
|
||||
if home, err := os.UserHomeDir(); err == nil && strings.TrimSpace(home) != "" {
|
||||
paths = append(paths, filepath.Join(home, "Applications", "ChatGPT.app"))
|
||||
}
|
||||
return &darwinProvider{
|
||||
appSessionID: uuid.NewString(),
|
||||
appPaths: paths,
|
||||
}
|
||||
}
|
||||
|
||||
func (p *darwinProvider) Check(ctx context.Context) error {
|
||||
checkCtx, cancel := context.WithTimeout(ctx, attestationTimeout)
|
||||
defer cancel()
|
||||
_, _, _, err := p.resolveRuntime(checkCtx)
|
||||
return err
|
||||
}
|
||||
|
||||
func (p *darwinProvider) resolveRuntime(ctx context.Context) (string, string, string, error) {
|
||||
if runtime.GOARCH != "arm64" {
|
||||
return "", "", "", errors.New("live attestation currently requires Apple Silicon; Intel macOS is not supported")
|
||||
}
|
||||
appPath, err := p.findApplication()
|
||||
if err != nil {
|
||||
return "", "", "", err
|
||||
}
|
||||
resourcesPath := filepath.Join(appPath, "Contents", "Resources")
|
||||
nodePath := filepath.Join(resourcesPath, "cua_node", "bin", "node")
|
||||
modulePath := filepath.Join(resourcesPath, "native", "devicecheck.node")
|
||||
for filePath, label := range map[string]string{
|
||||
nodePath: "bundled Node.js runtime",
|
||||
modulePath: "DeviceCheck native module",
|
||||
} {
|
||||
if info, statErr := os.Stat(filePath); statErr != nil || info.IsDir() {
|
||||
return "", "", "", fmt.Errorf("%w: ChatGPT app is missing its %s", ErrChatGPTAppMissing, label)
|
||||
}
|
||||
}
|
||||
bundleID, err := readBundleIdentifier(ctx, appPath)
|
||||
if err != nil {
|
||||
return "", "", "", err
|
||||
}
|
||||
return nodePath, modulePath, bundleID, nil
|
||||
}
|
||||
|
||||
func (p *darwinProvider) Generate(ctx context.Context) (string, error) {
|
||||
runCtx, cancel := context.WithTimeout(ctx, attestationTimeout)
|
||||
defer cancel()
|
||||
nodePath, modulePath, bundleID, err := p.resolveRuntime(runCtx)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
signals, err := p.readSignals(runCtx)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
signalsJSON, err := json.Marshal(signals)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("encode Live attestation signals: %w", err)
|
||||
}
|
||||
|
||||
command := exec.CommandContext(runCtx, nodePath, "-e", deviceCheckScript)
|
||||
command.Env = []string{
|
||||
"PATH=/usr/bin:/bin",
|
||||
"SUB2API_DEVICECHECK_MODULE=" + modulePath,
|
||||
"SUB2API_ATTESTATION_BUNDLE_ID=" + bundleID,
|
||||
"SUB2API_ATTESTATION_SIGNALS=" + string(signalsJSON),
|
||||
}
|
||||
var stdout bytes.Buffer
|
||||
var stderr bytes.Buffer
|
||||
command.Stdout = &stdout
|
||||
command.Stderr = &stderr
|
||||
if err := command.Run(); err != nil {
|
||||
if errors.Is(runCtx.Err(), context.DeadlineExceeded) {
|
||||
return "", errors.New("ChatGPT DeviceCheck token generation timed out")
|
||||
}
|
||||
reason := strings.TrimSpace(stderr.String())
|
||||
if len(reason) > 240 {
|
||||
reason = reason[:240]
|
||||
}
|
||||
if reason == "" {
|
||||
reason = err.Error()
|
||||
}
|
||||
return "", fmt.Errorf("ChatGPT DeviceCheck token generation failed: %s", reason)
|
||||
}
|
||||
header := strings.TrimSpace(stdout.String())
|
||||
if len(header) < 20 || len(header) > 16*1024 || !json.Valid([]byte(header)) {
|
||||
return "", errors.New("ChatGPT DeviceCheck returned a malformed attestation")
|
||||
}
|
||||
return header, nil
|
||||
}
|
||||
|
||||
func (p *darwinProvider) findApplication() (string, error) {
|
||||
for _, appPath := range p.appPaths {
|
||||
info, err := os.Stat(appPath)
|
||||
if err == nil && info.IsDir() {
|
||||
return appPath, nil
|
||||
}
|
||||
}
|
||||
return "", ErrChatGPTAppMissing
|
||||
}
|
||||
|
||||
func readBundleIdentifier(ctx context.Context, appPath string) (string, error) {
|
||||
infoPlist := filepath.Join(appPath, "Contents", "Info.plist")
|
||||
output, err := exec.CommandContext(
|
||||
ctx,
|
||||
"/usr/bin/plutil",
|
||||
"-extract",
|
||||
"CFBundleIdentifier",
|
||||
"raw",
|
||||
infoPlist,
|
||||
).Output()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%w: cannot read its bundle identifier", ErrChatGPTAppMissing)
|
||||
}
|
||||
bundleID := strings.TrimSpace(string(output))
|
||||
if !strings.HasPrefix(bundleID, "com.openai.") {
|
||||
return "", errors.New("the installed ChatGPT app has an unexpected bundle identifier")
|
||||
}
|
||||
return bundleID, nil
|
||||
}
|
||||
|
||||
func (p *darwinProvider) readSignals(ctx context.Context) (deviceSignals, error) {
|
||||
const script = `ObjC.import("Foundation"); ObjC.import("AppKit");
|
||||
const screen = $.NSScreen.mainScreen;
|
||||
const frame = screen.frame;
|
||||
JSON.stringify({
|
||||
locale: ObjC.unwrap($.NSLocale.currentLocale.localeIdentifier),
|
||||
languages: ObjC.deepUnwrap($.NSLocale.preferredLanguages),
|
||||
timezone: ObjC.unwrap($.NSTimeZone.localTimeZone.name),
|
||||
width: Number(frame.size.width),
|
||||
height: Number(frame.size.height),
|
||||
scale: Number(screen.backingScaleFactor)
|
||||
})`
|
||||
output, err := exec.CommandContext(ctx, "/usr/bin/osascript", "-l", "JavaScript", "-e", script).Output()
|
||||
if err != nil {
|
||||
return deviceSignals{}, fmt.Errorf("read macOS signals for Live attestation: %w", err)
|
||||
}
|
||||
var values macOSSignals
|
||||
if err := json.Unmarshal(output, &values); err != nil {
|
||||
return deviceSignals{}, fmt.Errorf("decode macOS signals for Live attestation: %w", err)
|
||||
}
|
||||
locale := truncateSignal(values.Locale, 64, "unknown")
|
||||
languages := values.Languages
|
||||
if len(languages) == 0 {
|
||||
languages = []string{locale}
|
||||
}
|
||||
if len(languages) > 16 {
|
||||
languages = languages[:16]
|
||||
}
|
||||
for index := range languages {
|
||||
languages[index] = truncateSignal(languages[index], 64, locale)
|
||||
}
|
||||
scale := values.Scale
|
||||
if scale <= 0 {
|
||||
scale = 1
|
||||
}
|
||||
return deviceSignals{
|
||||
SchemaVersion: 1,
|
||||
PreferredLanguages: languages,
|
||||
Locale: locale,
|
||||
Timezone: truncateSignal(values.Timezone, 64, "unknown"),
|
||||
ScreenSizeSum: max(0, int(values.Width+values.Height+0.5)),
|
||||
ScreenScale: scale,
|
||||
AppSessionID: truncateSignal(p.appSessionID, 128, uuid.NewString()),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func truncateSignal(value string, limit int, fallback string) string {
|
||||
value = strings.TrimSpace(value)
|
||||
if value == "" {
|
||||
value = fallback
|
||||
}
|
||||
if len(value) > limit {
|
||||
return value[:limit]
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
const deviceCheckScript = `
|
||||
const addon = require(process.env.SUB2API_DEVICECHECK_MODULE);
|
||||
const signals = JSON.parse(process.env.SUB2API_ATTESTATION_SIGNALS);
|
||||
const bundleID = process.env.SUB2API_ATTESTATION_BUNDLE_ID;
|
||||
|
||||
function head(major, value) {
|
||||
if (value < 24) return Buffer.from([major + value]);
|
||||
if (value <= 255) return Buffer.from([major + 24, value]);
|
||||
if (value <= 65535) {
|
||||
const out = Buffer.allocUnsafe(3);
|
||||
out[0] = major + 25;
|
||||
out.writeUInt16BE(value, 1);
|
||||
return out;
|
||||
}
|
||||
const out = Buffer.allocUnsafe(5);
|
||||
out[0] = major + 26;
|
||||
out.writeUInt32BE(value, 1);
|
||||
return out;
|
||||
}
|
||||
function uint(value) { return head(0, value); }
|
||||
function text(value) {
|
||||
const body = Buffer.from(value, "utf8");
|
||||
return Buffer.concat([head(96, body.length), body]);
|
||||
}
|
||||
function float(value) {
|
||||
if (Number.isSafeInteger(value) && value >= 0) return uint(value);
|
||||
const out = Buffer.allocUnsafe(9);
|
||||
out[0] = 251;
|
||||
out.writeDoubleBE(value, 1);
|
||||
return out;
|
||||
}
|
||||
function array(values) { return Buffer.concat([head(128, values.length), ...values]); }
|
||||
function map(entries) {
|
||||
return Buffer.concat([head(160, entries.length), ...entries.flatMap(([key, value]) => [uint(key), value])]);
|
||||
}
|
||||
function field(key, value) { return Buffer.concat([text(key), text(value)]); }
|
||||
function base64url(value) {
|
||||
return value.toString("base64").replaceAll("+", "-").replaceAll("/", "_").replace(/=+$/u, "");
|
||||
}
|
||||
|
||||
(async () => {
|
||||
const result = await addon.generateToken();
|
||||
if (!result || !result.supported) throw new Error("DeviceCheck is not supported on this Mac");
|
||||
if (!result.tokenBase64) throw new Error("DeviceCheck returned no token");
|
||||
const fingerprint = map([
|
||||
[0, uint(signals.schemaVersion)],
|
||||
[1, array(signals.preferredLanguages.map(text))],
|
||||
[2, text(signals.locale)],
|
||||
[3, text(signals.timezone)],
|
||||
[4, uint(signals.screenSizeSum)],
|
||||
[5, float(signals.screenScale)],
|
||||
[6, text(signals.appSessionId)]
|
||||
]);
|
||||
const fields = [
|
||||
field("token", result.tokenBase64),
|
||||
field("bundle_id", bundleID),
|
||||
Buffer.concat([text("f"), head(64, fingerprint.length), fingerprint])
|
||||
];
|
||||
if (result.latencyMs != null) {
|
||||
fields.push(Buffer.concat([text("t"), float(result.latencyMs)]));
|
||||
}
|
||||
const token = "v1." + base64url(Buffer.concat([Buffer.from([160 + fields.length]), ...fields]));
|
||||
process.stdout.write(JSON.stringify({v: 1, s: 0, t: token}));
|
||||
})().catch((error) => {
|
||||
process.stderr.write(error instanceof Error ? error.message : String(error));
|
||||
process.exitCode = 1;
|
||||
});`
|
||||
@@ -0,0 +1,19 @@
|
||||
//go:build !darwin
|
||||
|
||||
package liveattestation
|
||||
|
||||
import "context"
|
||||
|
||||
type unsupportedProvider struct{}
|
||||
|
||||
func NewProvider() Provider {
|
||||
return unsupportedProvider{}
|
||||
}
|
||||
|
||||
func (unsupportedProvider) Check(context.Context) error {
|
||||
return ErrUnsupportedPlatform
|
||||
}
|
||||
|
||||
func (unsupportedProvider) Generate(context.Context) (string, error) {
|
||||
return "", ErrUnsupportedPlatform
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
//go:build !darwin
|
||||
|
||||
package liveattestation
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestUnsupportedProviderReturnsExplicitPlatformError(t *testing.T) {
|
||||
provider := NewProvider()
|
||||
if err := provider.Check(context.Background()); !errors.Is(err, ErrUnsupportedPlatform) {
|
||||
t.Fatalf("Check() error = %v, want ErrUnsupportedPlatform", err)
|
||||
}
|
||||
_, err := provider.Generate(context.Background())
|
||||
if !errors.Is(err, ErrUnsupportedPlatform) {
|
||||
t.Fatalf("Generate() error = %v, want ErrUnsupportedPlatform", err)
|
||||
}
|
||||
}
|
||||
@@ -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,73 @@
|
||||
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, ok := regular.(service.LiveConcurrencyCache)
|
||||
require.True(t, ok)
|
||||
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, ok := regular.(service.LiveConcurrencyCache)
|
||||
require.True(t, ok)
|
||||
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,142 @@ 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,
|
||||
"attestation": record.AttestationCiphertext,
|
||||
}
|
||||
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"],
|
||||
AttestationCiphertext: values["attestation"],
|
||||
}, 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,62 @@
|
||||
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, ok := NewGatewayCache(client).(service.LiveCallStore)
|
||||
require.True(t, ok)
|
||||
otherInstance, ok := NewGatewayCache(client).(service.LiveCallStore)
|
||||
require.True(t, ok)
|
||||
record := &service.LiveCallRecord{
|
||||
CallID: "call_secret",
|
||||
CallHash: HashLiveCallID("call_secret"),
|
||||
AccountID: 11,
|
||||
APIKeyID: 22,
|
||||
UserID: 33,
|
||||
GroupID: 44,
|
||||
LeaseID: "lease",
|
||||
Model: "gpt-live-test",
|
||||
AttestationCiphertext: "encrypted-attestation",
|
||||
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)
|
||||
require.Equal(t, record.AttestationCiphertext, loaded.AttestationCiphertext)
|
||||
|
||||
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,
|
||||
|
||||
@@ -317,6 +317,7 @@ func registerGroupRoutes(admin *gin.RouterGroup, h *handler.Handlers) {
|
||||
groups.GET("/all", h.Admin.Group.GetAll)
|
||||
groups.GET("/usage-summary", h.Admin.Group.GetUsageSummary)
|
||||
groups.GET("/capacity-summary", h.Admin.Group.GetCapacitySummary)
|
||||
groups.GET("/live-capability", h.Admin.Group.GetLiveCapability)
|
||||
groups.PUT("/sort-order", h.Admin.Group.UpdateSortOrder)
|
||||
groups.GET("/:id/models-list-candidates", h.Admin.Group.GetModelsListCandidates)
|
||||
groups.GET("/:id/composite-routes", h.Admin.Group.ListCompositeRoutes)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -19,6 +19,7 @@ import (
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/ip"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/openai"
|
||||
"github.com/Wei-Shaw/sub2api/internal/platform/liveattestation"
|
||||
"github.com/Wei-Shaw/sub2api/internal/util/responseheaders"
|
||||
"github.com/cespare/xxhash/v2"
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -410,6 +411,8 @@ type OpenAIGatewayService struct {
|
||||
balanceNotifyService *BalanceNotifyService
|
||||
settingService *SettingService
|
||||
userPlatformQuotaRepo UserPlatformQuotaRepository
|
||||
liveAttestation liveattestation.Provider
|
||||
liveAttestationCipher SecretEncryptor
|
||||
|
||||
openaiWSPoolOnce sync.Once
|
||||
openaiWSStateStoreOnce sync.Once
|
||||
@@ -499,6 +502,8 @@ func NewOpenAIGatewayService(
|
||||
balanceNotifyService: balanceNotifyService,
|
||||
settingService: settingService,
|
||||
userPlatformQuotaRepo: userPlatformQuotaRepo,
|
||||
liveAttestation: liveattestation.NewProvider(),
|
||||
liveAttestationCipher: newLiveAttestationCipher(cfg),
|
||||
responseHeaderFilter: compileResponseHeaderFilter(cfg),
|
||||
codexSnapshotThrottle: newAccountWriteThrottle(openAICodexSnapshotPersistMinInterval),
|
||||
openaiModelTransient: newOpenAIAccountModelTransientState(openAIModelTransientDefaultMax),
|
||||
|
||||
@@ -0,0 +1,792 @@
|
||||
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
|
||||
}
|
||||
attestation, attestationCiphertext, err := s.prepareLiveAttestation(ctx)
|
||||
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, attestation)
|
||||
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,
|
||||
AttestationCiphertext: attestationCiphertext,
|
||||
}
|
||||
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,
|
||||
attestation string,
|
||||
) (*LiveCallCreated, error) {
|
||||
token, _, err := s.GetAccessToken(ctx, account)
|
||||
if err != nil {
|
||||
logLiveCreateStageFailure(ctx, account.ID, "access_token", err)
|
||||
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 {
|
||||
logLiveCreateStageFailure(ctx, account.ID, "authentication_headers", err)
|
||||
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 {
|
||||
logLiveCreateStageFailure(ctx, account.ID, "account_headers", err)
|
||||
return nil, err
|
||||
}
|
||||
upstreamReq.Header.Set("Content-Type", "application/json")
|
||||
upstreamReq.Header.Set("Accept", "application/sdp")
|
||||
upstreamReq.Header.Set(liveAttestationHeader, attestation)
|
||||
applyLiveUpstreamIdentityHeaders(upstreamReq.Header)
|
||||
|
||||
resp, err := s.httpUpstream.Do(upstreamReq, resolveAccountProxyURL(account), account.ID, account.Concurrency)
|
||||
if err != nil {
|
||||
logLiveCreateStageFailure(ctx, account.ID, "upstream_transport", err)
|
||||
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 logLiveCreateStageFailure(ctx context.Context, accountID int64, stage string, err error) {
|
||||
logger.FromContext(ctx).Warn(
|
||||
"OpenAI Live 创建阶段失败",
|
||||
zap.Int64("account_id", accountID),
|
||||
zap.String("stage", stage),
|
||||
zap.String("error_type", fmt.Sprintf("%T", err)),
|
||||
)
|
||||
}
|
||||
|
||||
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)
|
||||
if strings.TrimSpace(headers.Get("session-id")) == "" {
|
||||
headers.Set("session-id", uuid.NewString())
|
||||
}
|
||||
if strings.TrimSpace(headers.Get("thread-id")) == "" {
|
||||
headers.Set("thread-id", uuid.NewString())
|
||||
}
|
||||
// Realtime/Live 不使用 Responses 的实验头。
|
||||
headers.Del("OpenAI-Beta")
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) liveSidebandHeaders(
|
||||
ctx context.Context,
|
||||
account *Account,
|
||||
record *LiveCallRecord,
|
||||
) (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
|
||||
}
|
||||
attestation, err := s.decryptLiveAttestation(record)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
headers.Set(liveAttestationHeader, attestation)
|
||||
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, record)
|
||||
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 func() { _ = 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 liveSessionEnded(runErr) || !time.Now().Before(record.ExpiresAt) {
|
||||
s.finalizeLiveCall(record)
|
||||
return runErr
|
||||
}
|
||||
go s.observeLiveCall(record.CallHash)
|
||||
return runErr
|
||||
}
|
||||
|
||||
// liveSessionEnded 判断控制连接的退出原因是否意味着会话已终结(应 finalize:写
|
||||
// usage log 并释放租约),而不是可以交给 observer 重连的临时错误。
|
||||
//
|
||||
// ErrLiveUnavailable 在控制循环里只会来自租约续租失败。RefreshLiveLease 的 Lua 在
|
||||
// leaseID 被 GC 后不会重新写入,重连也拿不回并发槽 —— 若按临时错误重试,会话会以
|
||||
// 约 1 秒一轮的节奏空转到 ExpiresAt,期间持着上游连接却不计入任何并发限制。
|
||||
func liveSessionEnded(err error) bool {
|
||||
return errors.Is(err, ErrLiveCallNotFound) ||
|
||||
errors.Is(err, ErrLiveUnavailable) ||
|
||||
errors.Is(err, context.DeadlineExceeded)
|
||||
}
|
||||
|
||||
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 liveSessionEnded(runErr) {
|
||||
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)
|
||||
// 过期不在此处判定:返回 true 让调用方回到循环顶部的过期分支,由它 finalize
|
||||
// (写 usage log + 释放租约)。在这里直接返回 false 会让会话静默结束、不留记录。
|
||||
return err == nil && controller == LiveControllerObserver
|
||||
}
|
||||
|
||||
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,110 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"strings"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
)
|
||||
|
||||
const liveAttestationHeader = "x-oai-attestation"
|
||||
|
||||
type liveAttestationAES struct {
|
||||
key [32]byte
|
||||
}
|
||||
|
||||
func newLiveAttestationCipher(cfg *config.Config) SecretEncryptor {
|
||||
if cfg == nil || strings.TrimSpace(cfg.JWT.Secret) == "" {
|
||||
return nil
|
||||
}
|
||||
return &liveAttestationAES{
|
||||
key: sha256.Sum256([]byte("sub2api/live-attestation/v1\x00" + cfg.JWT.Secret)),
|
||||
}
|
||||
}
|
||||
|
||||
func (c *liveAttestationAES) Encrypt(plaintext string) (string, error) {
|
||||
block, err := aes.NewCipher(c.key[:])
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
gcm, err := cipher.NewGCM(block)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
nonce := make([]byte, gcm.NonceSize())
|
||||
if _, err := io.ReadFull(rand.Reader, nonce); err != nil {
|
||||
return "", fmt.Errorf("generate Live attestation nonce: %w", err)
|
||||
}
|
||||
encrypted := gcm.Seal(nonce, nonce, []byte(plaintext), nil)
|
||||
return base64.RawStdEncoding.EncodeToString(encrypted), nil
|
||||
}
|
||||
|
||||
func (c *liveAttestationAES) Decrypt(ciphertext string) (string, error) {
|
||||
encrypted, err := base64.RawStdEncoding.DecodeString(ciphertext)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("decode Live attestation: %w", err)
|
||||
}
|
||||
block, err := aes.NewCipher(c.key[:])
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
gcm, err := cipher.NewGCM(block)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if len(encrypted) < gcm.NonceSize() {
|
||||
return "", errors.New("encrypted Live attestation is too short")
|
||||
}
|
||||
plaintext, err := gcm.Open(nil, encrypted[:gcm.NonceSize()], encrypted[gcm.NonceSize():], nil)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("decrypt Live attestation: %w", err)
|
||||
}
|
||||
return string(plaintext), nil
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) prepareLiveAttestation(ctx context.Context) (string, string, error) {
|
||||
if s == nil || s.liveAttestation == nil {
|
||||
return "", "", &LiveAttestationUnavailableError{
|
||||
Reason: "Sub2API has no platform DeviceCheck provider",
|
||||
}
|
||||
}
|
||||
if s.liveAttestationCipher == nil {
|
||||
return "", "", &LiveAttestationUnavailableError{
|
||||
Reason: "JWT secret is required to protect the Sideband attestation",
|
||||
}
|
||||
}
|
||||
header, err := s.liveAttestation.Generate(ctx)
|
||||
if err != nil {
|
||||
return "", "", &LiveAttestationUnavailableError{Reason: err.Error()}
|
||||
}
|
||||
ciphertext, err := s.liveAttestationCipher.Encrypt(header)
|
||||
if err != nil {
|
||||
return "", "", &LiveAttestationUnavailableError{
|
||||
Reason: "failed to protect the generated DeviceCheck attestation",
|
||||
}
|
||||
}
|
||||
return header, ciphertext, nil
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) decryptLiveAttestation(record *LiveCallRecord) (string, error) {
|
||||
if record == nil || strings.TrimSpace(record.AttestationCiphertext) == "" || s.liveAttestationCipher == nil {
|
||||
return "", &LiveAttestationUnavailableError{
|
||||
Reason: "the Live call has no reusable DeviceCheck attestation",
|
||||
}
|
||||
}
|
||||
header, err := s.liveAttestationCipher.Decrypt(record.AttestationCiphertext)
|
||||
if err != nil {
|
||||
return "", &LiveAttestationUnavailableError{
|
||||
Reason: "the Live call DeviceCheck attestation cannot be decrypted on this instance",
|
||||
}
|
||||
}
|
||||
return header, nil
|
||||
}
|
||||
@@ -0,0 +1,469 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
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,
|
||||
}
|
||||
attestationCipher := newLiveAttestationCipher(&config.Config{
|
||||
JWT: config.JWTConfig{Secret: "live-sideband-test-secret"},
|
||||
})
|
||||
var err error
|
||||
record.AttestationCiphertext, err = attestationCipher.Encrypt(`{"v":1,"s":0,"t":"v1.sideband"}`)
|
||||
require.NoError(t, err)
|
||||
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,
|
||||
liveAttestationCipher: attestationCipher,
|
||||
}
|
||||
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 func() { _ = 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 func() { _ = 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"))
|
||||
require.Equal(t, `{"v":1,"s":0,"t":"v1.sideband"}`, dialer.headers.Get(liveAttestationHeader))
|
||||
upstream.reads <- liveTestFrame{err: coderws.CloseError{Code: coderws.StatusNormalClosure}}
|
||||
require.ErrorIs(t, <-proxyResult, ErrLiveCallNotFound)
|
||||
}
|
||||
|
||||
// TestLiveSessionEndedTreatsLeaseLossAsTerminal 锁定:租约续租失败(ErrLiveUnavailable)
|
||||
// 必须判为会话终结。RefreshLiveLease 的 Lua 在 leaseID 被 GC 后不会重新写入,若把它
|
||||
// 当临时错误交给 observer 重连,会话会空转到 ExpiresAt 且不计入任何并发限制。
|
||||
func TestLiveSessionEndedTreatsLeaseLossAsTerminal(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
err error
|
||||
want bool
|
||||
}{
|
||||
{"租约丢失", ErrLiveUnavailable, true},
|
||||
{"租约丢失(被包装)", fmt.Errorf("refresh live lease: %w", ErrLiveUnavailable), true},
|
||||
{"上游报告会话已关闭", ErrLiveCallNotFound, true},
|
||||
{"到达会话时长上限", context.DeadlineExceeded, true},
|
||||
{"控制权被他人接管", ErrLiveControllerChanged, false},
|
||||
{"临时读错误", errors.New("unexpected EOF"), false},
|
||||
{"无错误", nil, false},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
require.Equal(t, tc.want, liveSessionEnded(tc.err))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestWaitForLiveObserverRetryLeavesExpiryToLoopFinalize 锁定:已过期但控制权仍在
|
||||
// observer 手上时返回 true,让调用方回到 observeLiveCall 循环顶部的过期分支去
|
||||
// finalize(写 usage log + 释放租约)。在此处直接返回 false 会让会话静默结束、不留记录。
|
||||
func TestWaitForLiveObserverRetryLeavesExpiryToLoopFinalize(t *testing.T) {
|
||||
record := &LiveCallRecord{
|
||||
CallID: "call_expired",
|
||||
CallHash: hashLiveCallID("call_expired"),
|
||||
Controller: LiveControllerObserver,
|
||||
ExpiresAt: time.Now().Add(-time.Minute),
|
||||
}
|
||||
store := &liveTestStore{}
|
||||
require.NoError(t, store.SaveLiveCall(context.Background(), record, time.Hour))
|
||||
svc := &OpenAIGatewayService{cache: store}
|
||||
|
||||
require.True(t, svc.waitForLiveObserverRetry(record),
|
||||
"过期判定必须留给循环顶部,否则不会写 usage log")
|
||||
|
||||
// 控制权已被他人接管时仍必须停止重试,避免与新控制者抢同一个 call。
|
||||
require.NoError(t, store.SaveLiveCall(context.Background(), &LiveCallRecord{
|
||||
CallID: record.CallID,
|
||||
CallHash: record.CallHash,
|
||||
Controller: LiveControllerProxy,
|
||||
ExpiresAt: time.Now().Add(time.Hour),
|
||||
}, time.Hour))
|
||||
require.False(t, svc.waitForLiveObserverRetry(record))
|
||||
}
|
||||
@@ -0,0 +1,242 @@
|
||||
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
|
||||
}
|
||||
|
||||
type liveAttestationStub struct {
|
||||
header string
|
||||
err error
|
||||
}
|
||||
|
||||
func (s liveAttestationStub) Check(context.Context) error {
|
||||
return s.err
|
||||
}
|
||||
|
||||
func (s liveAttestationStub) Generate(context.Context) (string, error) {
|
||||
return s.header, s.err
|
||||
}
|
||||
|
||||
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,
|
||||
}, `{"v":1,"s":0,"t":"v1.test"}`)
|
||||
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.Equal(t, `{"v":1,"s":0,"t":"v1.test"}`, upstream.request.Header.Get(liveAttestationHeader))
|
||||
require.NotEmpty(t, upstream.request.Header.Get("Session-Id"))
|
||||
require.NotEmpty(t, upstream.request.Header.Get("Thread-Id"))
|
||||
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 TestLiveAttestationCipherRoundTripAndRejectsOtherInstanceKey(t *testing.T) {
|
||||
first := newLiveAttestationCipher(&config.Config{
|
||||
JWT: config.JWTConfig{Secret: "first-live-secret"},
|
||||
})
|
||||
second := newLiveAttestationCipher(&config.Config{
|
||||
JWT: config.JWTConfig{Secret: "second-live-secret"},
|
||||
})
|
||||
require.NotNil(t, first)
|
||||
require.NotNil(t, second)
|
||||
|
||||
ciphertext, err := first.Encrypt(`{"v":1,"s":0,"t":"v1.opaque"}`)
|
||||
require.NoError(t, err)
|
||||
require.NotContains(t, ciphertext, "opaque")
|
||||
|
||||
plaintext, err := first.Decrypt(ciphertext)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, `{"v":1,"s":0,"t":"v1.opaque"}`, plaintext)
|
||||
|
||||
_, err = second.Decrypt(ciphertext)
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
func TestPrepareLiveAttestationEncryptsHeaderAndReturnsExplicitProviderError(t *testing.T) {
|
||||
cipher := newLiveAttestationCipher(&config.Config{
|
||||
JWT: config.JWTConfig{Secret: "live-attestation-test-secret"},
|
||||
})
|
||||
service := &OpenAIGatewayService{
|
||||
liveAttestation: liveAttestationStub{header: `{"v":1,"s":0,"t":"v1.test"}`},
|
||||
liveAttestationCipher: cipher,
|
||||
}
|
||||
header, ciphertext, err := service.prepareLiveAttestation(context.Background())
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, `{"v":1,"s":0,"t":"v1.test"}`, header)
|
||||
require.NotContains(t, ciphertext, "v1.test")
|
||||
decrypted, err := cipher.Decrypt(ciphertext)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, header, decrypted)
|
||||
|
||||
service.liveAttestation = liveAttestationStub{err: errors.New("macOS app missing")}
|
||||
_, _, err = service.prepareLiveAttestation(context.Background())
|
||||
var unavailable *LiveAttestationUnavailableError
|
||||
require.ErrorAs(t, err, &unavailable)
|
||||
require.Contains(t, unavailable.Error(), "macOS app missing")
|
||||
}
|
||||
|
||||
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,103 @@
|
||||
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")
|
||||
)
|
||||
|
||||
type LiveAttestationUnavailableError struct {
|
||||
Reason string
|
||||
}
|
||||
|
||||
func (e *LiveAttestationUnavailableError) Error() string {
|
||||
if e == nil || e.Reason == "" {
|
||||
return "Live attestation is unavailable"
|
||||
}
|
||||
return "Live attestation is unavailable: " + e.Reason
|
||||
}
|
||||
|
||||
// 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
|
||||
// AttestationCiphertext 仅用于让同一会话的 Sideband 复用创建时的证明。
|
||||
AttestationCiphertext 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 行为。
|
||||
|
||||
@@ -16,6 +16,11 @@ import type {
|
||||
PaginatedResponse
|
||||
} from '@/types'
|
||||
|
||||
export interface LiveCapability {
|
||||
supported: boolean
|
||||
reason?: string
|
||||
}
|
||||
|
||||
/**
|
||||
* List all groups with pagination
|
||||
* @param page - Page number (default: 1)
|
||||
@@ -81,6 +86,12 @@ export async function getByPlatform(platform: GroupPlatform): Promise<AdminGroup
|
||||
return getAll(platform)
|
||||
}
|
||||
|
||||
/** 获取当前 Sub2API 服务端的 Live 运行环境能力。 */
|
||||
export async function getLiveCapability(): Promise<LiveCapability> {
|
||||
const { data } = await apiClient.get<LiveCapability>('/admin/groups/live-capability')
|
||||
return data
|
||||
}
|
||||
|
||||
/**
|
||||
* Get group by ID
|
||||
* @param id - Group ID
|
||||
@@ -467,6 +478,7 @@ export const groupsAPI = {
|
||||
getAll,
|
||||
getByPlatform,
|
||||
getAllIncludingInactive,
|
||||
getLiveCapability,
|
||||
getById,
|
||||
getModelsListCandidates,
|
||||
create,
|
||||
|
||||
@@ -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,14 @@ 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. The Sub2API server must run on Apple Silicon macOS with the official ChatGPT app installed; client platforms are unrestricted.',
|
||||
unsupportedTitle: 'Current server does not support Live',
|
||||
unsupportedMessage: 'This Sub2API server cannot generate the required Live attestation. Live will not work even if enabled. Continue anyway?',
|
||||
enableAnyway: 'Enable anyway'
|
||||
},
|
||||
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,14 @@ export default {
|
||||
targetModelPlaceholder: '例如: gpt-5.4',
|
||||
removeExactMapping: '删除精确映射'
|
||||
},
|
||||
openaiLive: {
|
||||
title: 'OpenAI Live',
|
||||
allow: '允许访问 Live',
|
||||
hint: '启用后,此 OpenAI 分组的 API Key 可以创建并控制 Live 语音会话。默认关闭。运行 Sub2API 的服务端必须是 Apple Silicon Mac,并安装官方 ChatGPT App;客户端平台不受限制。',
|
||||
unsupportedTitle: '当前服务端不支持 Live',
|
||||
unsupportedMessage: '当前 Sub2API 服务端无法生成 Live 所需的设备证明,即使开启也不能使用。是否仍然开启?',
|
||||
enableAnyway: '仍然开启'
|
||||
},
|
||||
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="toggleLive('create')"
|
||||
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="toggleLive('edit')"
|
||||
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'"
|
||||
@@ -3496,6 +3562,17 @@
|
||||
@cancel="showDeleteDialog = false"
|
||||
/>
|
||||
|
||||
<ConfirmDialog
|
||||
:show="showUnsupportedLiveConfirm"
|
||||
:title="t('admin.groups.openaiLive.unsupportedTitle')"
|
||||
:message="t('admin.groups.openaiLive.unsupportedMessage')"
|
||||
:confirm-text="t('admin.groups.openaiLive.enableAnyway')"
|
||||
:cancel-text="t('common.cancel')"
|
||||
:danger="true"
|
||||
@confirm="confirmUnsupportedLive"
|
||||
@cancel="cancelUnsupportedLive"
|
||||
/>
|
||||
|
||||
<!-- Sort Order Modal -->
|
||||
<BaseDialog
|
||||
:show="showSortModal"
|
||||
@@ -4433,6 +4510,15 @@ let abortController: AbortController | null = null;
|
||||
const showCreateModal = ref(false);
|
||||
const showEditModal = ref(false);
|
||||
const showDeleteDialog = ref(false);
|
||||
const pendingLiveForm = ref<"create" | "edit" | null>(null);
|
||||
const showUnsupportedLiveConfirm = computed(
|
||||
() => pendingLiveForm.value !== null,
|
||||
);
|
||||
const liveCapability = ref<{ supported: boolean; reason?: string } | null>(null);
|
||||
let liveCapabilityRequest: Promise<{
|
||||
supported: boolean;
|
||||
reason?: string;
|
||||
}> | null = null;
|
||||
const showSortModal = ref(false);
|
||||
const submitting = ref(false);
|
||||
const sortSubmitting = ref(false);
|
||||
@@ -4535,6 +4621,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 +4971,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,
|
||||
@@ -5081,6 +5169,44 @@ const deleteConfirmMessage = computed(() => {
|
||||
return t("admin.groups.deleteConfirm", { name: deletingGroup.value.name });
|
||||
});
|
||||
|
||||
const loadLiveCapability = async () => {
|
||||
if (liveCapability.value) return liveCapability.value;
|
||||
if (!liveCapabilityRequest) {
|
||||
liveCapabilityRequest = adminAPI.groups
|
||||
.getLiveCapability()
|
||||
.catch(() => ({ supported: false }))
|
||||
.finally(() => {
|
||||
liveCapabilityRequest = null;
|
||||
});
|
||||
}
|
||||
liveCapability.value = await liveCapabilityRequest;
|
||||
return liveCapability.value ?? { supported: false };
|
||||
};
|
||||
|
||||
const toggleLive = async (target: "create" | "edit") => {
|
||||
const form = target === "create" ? createForm : editForm;
|
||||
if (form.allow_live) {
|
||||
form.allow_live = false;
|
||||
return;
|
||||
}
|
||||
const capability = await loadLiveCapability();
|
||||
if (capability.supported) {
|
||||
form.allow_live = true;
|
||||
return;
|
||||
}
|
||||
pendingLiveForm.value = target;
|
||||
};
|
||||
|
||||
const confirmUnsupportedLive = () => {
|
||||
if (pendingLiveForm.value === "create") createForm.allow_live = true;
|
||||
if (pendingLiveForm.value === "edit") editForm.allow_live = true;
|
||||
pendingLiveForm.value = null;
|
||||
};
|
||||
|
||||
const cancelUnsupportedLive = () => {
|
||||
pendingLiveForm.value = null;
|
||||
};
|
||||
|
||||
const loadGroups = async () => {
|
||||
if (abortController) {
|
||||
abortController.abort();
|
||||
@@ -5287,6 +5413,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 +5601,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 +5658,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 +6063,7 @@ watch(
|
||||
}
|
||||
if (newVal !== "openai") {
|
||||
resetMessagesDispatchFormState(createForm);
|
||||
createForm.allow_live = false;
|
||||
}
|
||||
createForm.max_reasoning_effort = normalizeReasoningEffortForPlatform(
|
||||
newVal,
|
||||
@@ -5976,6 +6106,7 @@ watch(
|
||||
}
|
||||
if (newVal !== "openai") {
|
||||
resetMessagesDispatchFormState(editForm);
|
||||
editForm.allow_live = false;
|
||||
}
|
||||
editForm.max_reasoning_effort = normalizeReasoningEffortForPlatform(
|
||||
newVal,
|
||||
@@ -6020,6 +6151,7 @@ watch(
|
||||
}
|
||||
if (newVal !== 'openai') {
|
||||
editForm.allow_messages_dispatch = false
|
||||
editForm.allow_live = false
|
||||
editForm.default_mapped_model = ''
|
||||
}
|
||||
}
|
||||
@@ -6085,6 +6217,7 @@ const saveSortOrder = async () => {
|
||||
|
||||
onMounted(() => {
|
||||
loadGroups();
|
||||
void loadLiveCapability();
|
||||
loadModelsListCandidates("create", 0, createForm.platform);
|
||||
document.addEventListener("click", handleClickOutside);
|
||||
});
|
||||
|
||||
@@ -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