diff --git a/controller/metaproxy_provision.go b/controller/metaproxy_provision.go new file mode 100644 index 000000000000..363ff5285ee2 --- /dev/null +++ b/controller/metaproxy_provision.go @@ -0,0 +1,310 @@ +package controller + +import ( + "errors" + "fmt" + "math" + "net/http" + "net/url" + "regexp" + "strings" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/model" + "github.com/gin-gonic/gin" +) + +var ( + metaproxyRevisionPattern = regexp.MustCompile(`^[0-9a-f]{40}$`) + metaproxyDigestPattern = regexp.MustCompile(`^[0-9a-f]{64}$`) +) + +type metaproxyProvisionChannelRequest struct { + Type int `json:"type"` + Key string `json:"key"` + Name string `json:"name"` + BaseURL string `json:"base_url"` + Models string `json:"models"` + ModelMapping string `json:"model_mapping"` + Group string `json:"group"` + Priority int64 `json:"priority"` + Weight uint `json:"weight"` + Status int `json:"status"` + TestModel string `json:"test_model"` + HeaderOverride string `json:"header_override"` +} + +type metaproxyProvisionOptionsRequest struct { + ModelRatio string `json:"model_ratio"` + CompletionRatio string `json:"completion_ratio"` + CacheRatio string `json:"cache_ratio"` + GroupRatio string `json:"group_ratio"` + UserUsableGroups string `json:"user_usable_groups"` +} + +type metaproxyProvisionRequest struct { + Revision string `json:"revision"` + Digest string `json:"digest"` + Channels []metaproxyProvisionChannelRequest `json:"channels"` + Options metaproxyProvisionOptionsRequest `json:"options"` +} + +func parseRatioMap(name, raw string, allowZero bool) (map[string]float64, error) { + values := make(map[string]float64) + if err := common.Unmarshal([]byte(raw), &values); err != nil { + return nil, fmt.Errorf("%s must be a JSON object of numbers: %w", name, err) + } + for key, value := range values { + if strings.TrimSpace(key) == "" || math.IsNaN(value) || math.IsInf(value, 0) { + return nil, fmt.Errorf("%s contains an invalid value for %q", name, key) + } + if value < 0 || (!allowZero && value == 0) { + return nil, fmt.Errorf("%s must contain positive values: %q", name, key) + } + } + return values, nil +} + +func parseUserUsableGroups(raw string) (map[string]string, error) { + values := make(map[string]string) + if err := common.Unmarshal([]byte(raw), &values); err != nil { + return nil, fmt.Errorf("UserUsableGroups must be a JSON object of strings: %w", err) + } + for group, label := range values { + if strings.TrimSpace(group) == "" || strings.TrimSpace(label) == "" { + return nil, fmt.Errorf("UserUsableGroups contains an empty group or label") + } + } + return values, nil +} + +func validateOptionalJSONObject(name, raw string) error { + if raw == "" { + return nil + } + value := make(map[string]any) + if err := common.Unmarshal([]byte(raw), &value); err != nil { + return fmt.Errorf("%s must be a JSON object: %w", name, err) + } + return nil +} + +func commaValues(raw string) ([]string, error) { + parts := strings.Split(raw, ",") + seen := make(map[string]struct{}, len(parts)) + for index, part := range parts { + part = strings.TrimSpace(part) + if part == "" { + return nil, errors.New("contains an empty item") + } + if _, duplicate := seen[part]; duplicate { + return nil, fmt.Errorf("contains duplicate item %q", part) + } + seen[part] = struct{}{} + parts[index] = part + } + return parts, nil +} + +func validateProvisionBaseURL(raw string) error { + if raw == "" { + return nil + } + parsed, err := url.Parse(raw) + if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") { + return errors.New("base_url must be an absolute HTTP(S) URL") + } + if parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" { + return errors.New("base_url must not contain credentials, a query, or a fragment") + } + return nil +} + +func validateMetaproxyProvisionRequest( + request metaproxyProvisionRequest, + idempotencyKey string, + expectedDigest string, +) error { + if !metaproxyRevisionPattern.MatchString(request.Revision) { + return errors.New("revision must be a lowercase 40-character Git SHA") + } + if !metaproxyDigestPattern.MatchString(request.Digest) { + return errors.New("digest must be a lowercase SHA-256 value") + } + if idempotencyKey != request.Digest { + return errors.New("Idempotency-Key must equal the requested configuration digest") + } + if expectedDigest != model.MetaproxyProvisionNoDigest && + !metaproxyDigestPattern.MatchString(expectedDigest) { + return errors.New("If-Match must be 'none' or a lowercase SHA-256 value") + } + if len(request.Channels) > 256 { + return errors.New("channels exceeds the 256-channel limit") + } + + modelRatios, err := parseRatioMap("ModelRatio", request.Options.ModelRatio, false) + if err != nil { + return err + } + if _, err := parseRatioMap("CompletionRatio", request.Options.CompletionRatio, false); err != nil { + return err + } + if _, err := parseRatioMap("CacheRatio", request.Options.CacheRatio, true); err != nil { + return err + } + groupRatios, err := parseRatioMap("GroupRatio", request.Options.GroupRatio, false) + if err != nil { + return err + } + usableGroups, err := parseUserUsableGroups(request.Options.UserUsableGroups) + if err != nil { + return err + } + for group := range usableGroups { + if _, priced := groupRatios[group]; !priced { + return fmt.Errorf("offered group %q is missing from GroupRatio", group) + } + } + for group := range groupRatios { + if _, offered := usableGroups[group]; !offered { + return fmt.Errorf("priced group %q is missing from UserUsableGroups", group) + } + } + + names := make(map[string]struct{}, len(request.Channels)) + for index, channel := range request.Channels { + prefix := fmt.Sprintf("channels[%d]", index) + if channel.Type <= 0 { + return fmt.Errorf("%s.type must be positive", prefix) + } + if strings.TrimSpace(channel.Key) == "" { + return fmt.Errorf("%s.key must not be empty", prefix) + } + if strings.TrimSpace(channel.Name) == "" || len(channel.Name) > 255 { + return fmt.Errorf("%s.name must contain 1 to 255 characters", prefix) + } + if _, duplicate := names[channel.Name]; duplicate { + return fmt.Errorf("duplicate channel name %q", channel.Name) + } + names[channel.Name] = struct{}{} + if err := validateProvisionBaseURL(channel.BaseURL); err != nil { + return fmt.Errorf("%s.%w", prefix, err) + } + models, err := commaValues(channel.Models) + if err != nil { + return fmt.Errorf("%s.models %w", prefix, err) + } + groups, err := commaValues(channel.Group) + if err != nil { + return fmt.Errorf("%s.group %w", prefix, err) + } + if channel.Status < 1 || channel.Status > 3 { + return fmt.Errorf("%s.status must be 1, 2, or 3", prefix) + } + if err := validateOptionalJSONObject(prefix+".model_mapping", channel.ModelMapping); err != nil { + return err + } + if err := validateOptionalJSONObject(prefix+".header_override", channel.HeaderOverride); err != nil { + return err + } + if channel.Status != 1 { + continue + } + offered := false + for _, group := range groups { + _, groupOffered := usableGroups[group] + offered = offered || groupOffered + } + if !offered { + continue + } + for _, modelName := range models { + if _, priced := modelRatios[modelName]; !priced { + return fmt.Errorf("enabled model %q in an offered group is missing from ModelRatio", modelName) + } + } + } + return nil +} + +func toMetaproxyProvisionConfig(request metaproxyProvisionRequest) model.MetaproxyProvisionConfig { + channels := make([]model.MetaproxyProvisionChannel, 0, len(request.Channels)) + for _, channel := range request.Channels { + models, _ := commaValues(channel.Models) + groups, _ := commaValues(channel.Group) + channels = append(channels, model.MetaproxyProvisionChannel{ + Type: channel.Type, + Key: channel.Key, + Name: channel.Name, + BaseURL: channel.BaseURL, + Models: strings.Join(models, ","), + ModelMapping: channel.ModelMapping, + Group: strings.Join(groups, ","), + Priority: channel.Priority, + Weight: channel.Weight, + Status: channel.Status, + TestModel: channel.TestModel, + HeaderOverride: channel.HeaderOverride, + }) + } + return model.MetaproxyProvisionConfig{ + Revision: request.Revision, + Digest: request.Digest, + Channels: channels, + Options: model.MetaproxyProvisionOptions{ + ModelRatio: request.Options.ModelRatio, + CompletionRatio: request.Options.CompletionRatio, + CacheRatio: request.Options.CacheRatio, + GroupRatio: request.Options.GroupRatio, + UserUsableGroups: request.Options.UserUsableGroups, + }, + } +} + +func ApplyMetaproxyProvision(c *gin.Context) { + var request metaproxyProvisionRequest + if err := common.DecodeJson(c.Request.Body, &request); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": "invalid JSON body"}) + return + } + expectedDigest := strings.Trim(c.GetHeader("If-Match"), `"`) + if err := validateMetaproxyProvisionRequest( + request, + c.GetHeader("Idempotency-Key"), + expectedDigest, + ); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": err.Error()}) + return + } + + result, err := model.ApplyMetaproxyProvision(toMetaproxyProvisionConfig(request), expectedDigest) + if errors.Is(err, model.ErrMetaproxyProvisionRequiresMemoryCache) { + c.JSON(http.StatusPreconditionFailed, gin.H{"success": false, "message": err.Error()}) + return + } + if errors.Is(err, model.ErrMetaproxyProvisionConflict) { + c.JSON(http.StatusConflict, gin.H{"success": false, "message": err.Error()}) + return + } + if err != nil { + common.ApiError(c, err) + return + } + recordManageAudit(c, "metaproxy.provision.apply", map[string]interface{}{ + "revision": request.Revision, + "digest": request.Digest, + "channel_count": len(request.Channels), + "already_applied": result.AlreadyApplied, + }) + c.JSON(http.StatusOK, gin.H{ + "success": true, + "message": "", + "data": gin.H{ + "revision": request.Revision, + "digest": request.Digest, + "previous_digest": result.PreviousDigest, + "already_applied": result.AlreadyApplied, + "restart_required": result.RestartRequired, + }, + }) +} diff --git a/controller/metaproxy_provision_test.go b/controller/metaproxy_provision_test.go new file mode 100644 index 000000000000..5a1d18719fac --- /dev/null +++ b/controller/metaproxy_provision_test.go @@ -0,0 +1,131 @@ +package controller + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/QuantumNous/new-api/common" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func validMetaproxyProvisionRequest() metaproxyProvisionRequest { + return metaproxyProvisionRequest{ + Revision: "ba57fd0526c6ce9e9869c225ac9997d96b2bbdea", + Digest: "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", + Channels: []metaproxyProvisionChannelRequest{ + { + Type: 1, + Key: "secret", + Name: "Upstream [one]", + BaseURL: "https://upstream.example/v1", + Models: "model-one", + Group: "standard", + Priority: 10, + Weight: 100, + Status: 1, + TestModel: "model-one", + }, + }, + Options: metaproxyProvisionOptionsRequest{ + ModelRatio: `{"model-one":1}`, + CompletionRatio: `{"model-one":2}`, + CacheRatio: `{"model-one":0.1}`, + GroupRatio: `{"standard":1}`, + UserUsableGroups: `{"standard":"Standard"}`, + }, + } +} + +func TestValidateMetaproxyProvisionRequestAcceptsCompleteConfig(t *testing.T) { + request := validMetaproxyProvisionRequest() + require.NoError(t, validateMetaproxyProvisionRequest(request, request.Digest, "none")) +} + +func TestValidateMetaproxyProvisionRequestRejectsMismatchedIdempotencyKey(t *testing.T) { + request := validMetaproxyProvisionRequest() + err := validateMetaproxyProvisionRequest( + request, + "bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb", + "none", + ) + require.ErrorContains(t, err, "Idempotency-Key") +} + +func TestValidateMetaproxyProvisionRequestRejectsDuplicateChannelNames(t *testing.T) { + request := validMetaproxyProvisionRequest() + request.Channels = append(request.Channels, request.Channels[0]) + err := validateMetaproxyProvisionRequest(request, request.Digest, "none") + require.ErrorContains(t, err, "duplicate channel name") +} + +func TestValidateMetaproxyProvisionRequestRejectsEnabledModelWithoutPrice(t *testing.T) { + request := validMetaproxyProvisionRequest() + request.Options.ModelRatio = `{}` + err := validateMetaproxyProvisionRequest(request, request.Digest, "none") + require.ErrorContains(t, err, "model-one") + require.ErrorContains(t, err, "ModelRatio") +} + +func TestValidateMetaproxyProvisionRequestAllowsDisabledUnpricedModel(t *testing.T) { + request := validMetaproxyProvisionRequest() + request.Channels[0].Status = 2 + request.Options.ModelRatio = `{}` + require.NoError(t, validateMetaproxyProvisionRequest(request, request.Digest, "none")) +} + +func TestValidateMetaproxyProvisionRequestAllowsUnpricedModelInUnpublishedGroup(t *testing.T) { + request := validMetaproxyProvisionRequest() + request.Options.ModelRatio = `{}` + request.Options.GroupRatio = `{}` + request.Options.UserUsableGroups = `{}` + require.NoError(t, validateMetaproxyProvisionRequest(request, request.Digest, "none")) +} + +func TestValidateMetaproxyProvisionRequestRejectsMalformedRatioJson(t *testing.T) { + request := validMetaproxyProvisionRequest() + request.Options.GroupRatio = `{"standard":"free"}` + err := validateMetaproxyProvisionRequest(request, request.Digest, "none") + require.ErrorContains(t, err, "GroupRatio") +} + +func TestCommaValuesNormalizesWhitespace(t *testing.T) { + values, err := commaValues("model-a, model-b\t") + require.NoError(t, err) + require.Equal(t, []string{"model-a", "model-b"}, values) +} + +func TestToMetaproxyProvisionConfigStoresNormalizedLists(t *testing.T) { + request := validMetaproxyProvisionRequest() + request.Channels[0].Models = "model-one, model-two\t" + request.Channels[0].Group = "standard, archive\t" + + config := toMetaproxyProvisionConfig(request) + require.Equal(t, "model-one,model-two", config.Channels[0].Models) + require.Equal(t, "standard,archive", config.Channels[0].Group) +} + +func TestApplyMetaproxyProvisionRequiresMemoryCache(t *testing.T) { + previous := common.MemoryCacheEnabled + common.MemoryCacheEnabled = false + t.Cleanup(func() { common.MemoryCacheEnabled = previous }) + + body, err := json.Marshal(validMetaproxyProvisionRequest()) + require.NoError(t, err) + recorder := httptest.NewRecorder() + context, _ := gin.CreateTestContext(recorder) + context.Request = httptest.NewRequest(http.MethodPost, "/api/metaproxy/provision", bytes.NewReader(body)) + context.Request.Header.Set( + "Idempotency-Key", + "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", + ) + context.Request.Header.Set("If-Match", "none") + + ApplyMetaproxyProvision(context) + + require.Equal(t, http.StatusPreconditionFailed, recorder.Code) + require.Contains(t, recorder.Body.String(), "MEMORY_CACHE_ENABLED=true") +} diff --git a/controller/misc.go b/controller/misc.go index fb2029878747..b31f3af026c6 100644 --- a/controller/misc.go +++ b/controller/misc.go @@ -93,6 +93,8 @@ func GetStatus(c *gin.Context) { "password_login_enabled": common.PasswordLoginEnabled, "password_register_enabled": common.PasswordRegisterEnabled, "default_use_auto_group": setting.DefaultUseAutoGroup, + "metaproxy_provision_digest": common.OptionMap[model.MetaproxyProvisionDigestOption], + "metaproxy_provision_revision": common.OptionMap[model.MetaproxyProvisionRevisionOption], "usd_exchange_rate": operation_setting.USDExchangeRate, "price": operation_setting.Price, diff --git a/model/channel_cache.go b/model/channel_cache.go index 81923017d79c..310037cc4d9e 100644 --- a/model/channel_cache.go +++ b/model/channel_cache.go @@ -106,8 +106,13 @@ func InitChannelCache() { func SyncChannelCache(frequency int) { for { time.Sleep(time.Duration(frequency) * time.Second) - common.SysLog("syncing channels from database") - InitChannelCache() + if !RunMetaproxyProvisionSyncIfReady(func() { + common.SysLog("syncing channels from database") + InitChannelCache() + }) { + common.SysLog("skipping channel sync while a metaproxy provision restart is pending") + continue + } } } diff --git a/model/metaproxy_provision.go b/model/metaproxy_provision.go new file mode 100644 index 000000000000..5f14c5ee0ba0 --- /dev/null +++ b/model/metaproxy_provision.go @@ -0,0 +1,333 @@ +package model + +import ( + "errors" + "fmt" + "sync" + "sync/atomic" + + "github.com/QuantumNous/new-api/common" + "gorm.io/gorm" + "gorm.io/gorm/clause" +) + +const ( + MetaproxyProvisionManagedTag = "metaproxy-provision" + MetaproxyProvisionDigestOption = "MetaproxyProvisionDigest" + MetaproxyProvisionRevisionOption = "MetaproxyProvisionRevision" + MetaproxyProvisionNoDigest = "none" +) + +var ( + ErrMetaproxyProvisionConflict = errors.New("metaproxy provision revision conflict") + ErrMetaproxyProvisionRequiresMemoryCache = errors.New("metaproxy provision requires MEMORY_CACHE_ENABLED=true") + metaproxyProvisionLock sync.Mutex + metaproxyProvisionRuntimeLock sync.RWMutex + provisionRuntimeFrozen atomic.Bool +) + +type MetaproxyProvisionChannel struct { + Type int `json:"type"` + Key string `json:"key"` + Name string `json:"name"` + BaseURL string `json:"base_url"` + Models string `json:"models"` + ModelMapping string `json:"model_mapping"` + Group string `json:"group"` + Priority int64 `json:"priority"` + Weight uint `json:"weight"` + Status int `json:"status"` + TestModel string `json:"test_model"` + HeaderOverride string `json:"header_override"` +} + +type MetaproxyProvisionOptions struct { + ModelRatio string `json:"model_ratio"` + CompletionRatio string `json:"completion_ratio"` + CacheRatio string `json:"cache_ratio"` + GroupRatio string `json:"group_ratio"` + UserUsableGroups string `json:"user_usable_groups"` +} + +func (options MetaproxyProvisionOptions) orderedValues() []Option { + return []Option{ + {Key: "ModelRatio", Value: options.ModelRatio}, + {Key: "CompletionRatio", Value: options.CompletionRatio}, + {Key: "CacheRatio", Value: options.CacheRatio}, + {Key: "GroupRatio", Value: options.GroupRatio}, + {Key: "UserUsableGroups", Value: options.UserUsableGroups}, + } +} + +type MetaproxyProvisionConfig struct { + Revision string `json:"revision"` + Digest string `json:"digest"` + Channels []MetaproxyProvisionChannel `json:"channels"` + Options MetaproxyProvisionOptions `json:"options"` +} + +type MetaproxyProvisionResult struct { + AlreadyApplied bool + RestartRequired bool + PreviousDigest string +} + +func IsMetaproxyProvisionRuntimeFrozen() bool { + return provisionRuntimeFrozen.Load() +} + +func RunMetaproxyProvisionSyncIfReady(syncFn func()) bool { + metaproxyProvisionRuntimeLock.RLock() + defer metaproxyProvisionRuntimeLock.RUnlock() + if provisionRuntimeFrozen.Load() { + return false + } + syncFn() + return true +} + +func activeMetaproxyProvisionDigest() string { + common.OptionMapRWMutex.RLock() + defer common.OptionMapRWMutex.RUnlock() + if digest := common.OptionMap[MetaproxyProvisionDigestOption]; digest != "" { + return digest + } + return MetaproxyProvisionNoDigest +} + +func provisionOptionValue(tx *gorm.DB, key string) (string, error) { + var option Option + err := lockForUpdate(tx).First(&option, &Option{Key: key}).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return MetaproxyProvisionNoDigest, nil + } + if err != nil { + return "", err + } + if option.Value == "" { + return MetaproxyProvisionNoDigest, nil + } + return option.Value, nil +} + +func optionMatches(tx *gorm.DB, key, value string) (bool, error) { + var option Option + err := tx.First(&option, &Option{Key: key}).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return false, nil + } + if err != nil { + return false, err + } + return option.Value == value, nil +} + +func channelMatches(current Channel, wanted MetaproxyProvisionChannel) bool { + return current.Type == wanted.Type && + current.Key == wanted.Key && + current.Name == wanted.Name && + valueOrEmpty(current.BaseURL) == wanted.BaseURL && + current.Models == wanted.Models && + current.GetModelMapping() == wanted.ModelMapping && + current.Group == wanted.Group && + current.GetPriority() == wanted.Priority && + current.GetWeight() == int(wanted.Weight) && + current.Status == wanted.Status && + valueOrEmpty(current.TestModel) == wanted.TestModel && + valueOrEmpty(current.HeaderOverride) == wanted.HeaderOverride +} + +func valueOrEmpty(value *string) string { + if value == nil { + return "" + } + return *value +} + +func desiredProvisionAlreadyStored(tx *gorm.DB, config MetaproxyProvisionConfig) (bool, error) { + var current []Channel + if err := tx.Where("tag = ?", MetaproxyProvisionManagedTag).Find(¤t).Error; err != nil { + return false, err + } + if len(current) != len(config.Channels) { + return false, nil + } + byName := make(map[string]Channel, len(current)) + for _, channel := range current { + if _, duplicate := byName[channel.Name]; duplicate { + return false, fmt.Errorf("duplicate managed channel name %q", channel.Name) + } + byName[channel.Name] = channel + } + for _, wanted := range config.Channels { + currentChannel, ok := byName[wanted.Name] + if !ok || !channelMatches(currentChannel, wanted) { + return false, nil + } + } + for _, option := range config.Options.orderedValues() { + matches, err := optionMatches(tx, option.Key, option.Value) + if err != nil || !matches { + return false, err + } + } + for _, option := range []Option{ + {Key: MetaproxyProvisionDigestOption, Value: config.Digest}, + {Key: MetaproxyProvisionRevisionOption, Value: config.Revision}, + } { + matches, err := optionMatches(tx, option.Key, option.Value) + if err != nil || !matches { + return false, err + } + } + return true, nil +} + +func pointer[T any](value T) *T { + return &value +} + +func provisionChannel(wanted MetaproxyProvisionChannel) Channel { + tag := MetaproxyProvisionManagedTag + return Channel{ + Type: wanted.Type, + Key: wanted.Key, + Name: wanted.Name, + BaseURL: pointer(wanted.BaseURL), + Models: wanted.Models, + ModelMapping: pointer(wanted.ModelMapping), + Group: wanted.Group, + Priority: pointer(wanted.Priority), + Weight: pointer(wanted.Weight), + Status: wanted.Status, + TestModel: pointer(wanted.TestModel), + HeaderOverride: pointer(wanted.HeaderOverride), + Tag: &tag, + } +} + +func saveProvisionOption(tx *gorm.DB, option Option) error { + return tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "key"}}, + DoUpdates: clause.AssignmentColumns([]string{"value"}), + }).Create(&option).Error +} + +func replaceManagedProvisionState(tx *gorm.DB, config MetaproxyProvisionConfig) error { + var current []Channel + if err := tx.Where("tag = ?", MetaproxyProvisionManagedTag).Find(¤t).Error; err != nil { + return err + } + byName := make(map[string]Channel, len(current)) + for _, channel := range current { + if _, duplicate := byName[channel.Name]; duplicate { + return fmt.Errorf("duplicate managed channel name %q", channel.Name) + } + byName[channel.Name] = channel + } + + wantedNames := make(map[string]struct{}, len(config.Channels)) + for _, wanted := range config.Channels { + wantedNames[wanted.Name] = struct{}{} + desired := provisionChannel(wanted) + if existing, ok := byName[wanted.Name]; ok { + desired.Id = existing.Id + desired.CreatedTime = existing.CreatedTime + desired.UsedQuota = existing.UsedQuota + desired.ChannelInfo = existing.ChannelInfo + if err := tx.Model(&Channel{Id: existing.Id}).Select( + "type", "key", "name", "base_url", "models", "model_mapping", "group", + "priority", "weight", "status", "test_model", "header_override", "tag", "channel_info", + ).Updates(&desired).Error; err != nil { + return err + } + if err := desired.UpdateAbilities(tx); err != nil { + return err + } + continue + } + desired.CreatedTime = common.GetTimestamp() + if err := tx.Create(&desired).Error; err != nil { + return err + } + if err := desired.AddAbilities(tx); err != nil { + return err + } + } + + for _, channel := range current { + if _, keep := wantedNames[channel.Name]; keep { + continue + } + if err := tx.Where("channel_id = ?", channel.Id).Delete(&Ability{}).Error; err != nil { + return err + } + if err := tx.Delete(&Channel{}, channel.Id).Error; err != nil { + return err + } + } + + for _, option := range config.Options.orderedValues() { + if err := saveProvisionOption(tx, option); err != nil { + return err + } + } + for _, option := range []Option{ + {Key: MetaproxyProvisionRevisionOption, Value: config.Revision}, + {Key: MetaproxyProvisionDigestOption, Value: config.Digest}, + } { + if err := saveProvisionOption(tx, option); err != nil { + return err + } + } + return nil +} + +func ApplyMetaproxyProvision( + config MetaproxyProvisionConfig, + expectedDigest string, +) (MetaproxyProvisionResult, error) { + if !common.MemoryCacheEnabled { + return MetaproxyProvisionResult{}, ErrMetaproxyProvisionRequiresMemoryCache + } + metaproxyProvisionLock.Lock() + defer metaproxyProvisionLock.Unlock() + metaproxyProvisionRuntimeLock.Lock() + defer metaproxyProvisionRuntimeLock.Unlock() + + result := MetaproxyProvisionResult{} + err := DB.Transaction(func(tx *gorm.DB) error { + currentDigest, err := provisionOptionValue(tx, MetaproxyProvisionDigestOption) + if err != nil { + return err + } + result.PreviousDigest = currentDigest + if currentDigest != expectedDigest && currentDigest != config.Digest { + return fmt.Errorf( + "%w: expected %q, current %q", + ErrMetaproxyProvisionConflict, + expectedDigest, + currentDigest, + ) + } + + stored, err := desiredProvisionAlreadyStored(tx, config) + if err != nil { + return err + } + if stored { + result.AlreadyApplied = true + return nil + } + return replaceManagedProvisionState(tx, config) + }) + if err != nil { + return MetaproxyProvisionResult{}, err + } + + result.RestartRequired = activeMetaproxyProvisionDigest() != config.Digest + if result.RestartRequired { + provisionRuntimeFrozen.Store(true) + } + return result, nil +} diff --git a/model/metaproxy_provision_test.go b/model/metaproxy_provision_test.go new file mode 100644 index 000000000000..9fe410d7aead --- /dev/null +++ b/model/metaproxy_provision_test.go @@ -0,0 +1,233 @@ +package model + +import ( + "errors" + "fmt" + "testing" + + "github.com/QuantumNous/new-api/common" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func setupMetaproxyProvisionTestDB(t *testing.T) *gorm.DB { + t.Helper() + dsn := fmt.Sprintf("file:metaproxy-provision-%s?mode=memory&cache=shared", t.Name()) + db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{}) + require.NoError(t, err) + require.NoError(t, db.AutoMigrate(&Channel{}, &Ability{}, &Option{})) + + previousDB := DB + previousMemoryCacheEnabled := common.MemoryCacheEnabled + DB = db + common.MemoryCacheEnabled = true + provisionRuntimeFrozen.Store(false) + common.OptionMapRWMutex.Lock() + previousOptions := common.OptionMap + common.OptionMap = map[string]string{ + MetaproxyProvisionDigestOption: "old-digest", + MetaproxyProvisionRevisionOption: "old-revision", + } + common.OptionMapRWMutex.Unlock() + t.Cleanup(func() { + DB = previousDB + common.MemoryCacheEnabled = previousMemoryCacheEnabled + provisionRuntimeFrozen.Store(false) + common.OptionMapRWMutex.Lock() + common.OptionMap = previousOptions + common.OptionMapRWMutex.Unlock() + }) + return db +} + +func provisionTestChannel(name, key, models string) MetaproxyProvisionChannel { + return MetaproxyProvisionChannel{ + Type: 1, + Key: key, + Name: name, + BaseURL: "https://upstream.example/v1", + Models: models, + Group: "standard", + Priority: 10, + Weight: 100, + Status: 1, + TestModel: models, + ModelMapping: "", + } +} + +func seedProvisionState(t *testing.T, db *gorm.DB) Channel { + t.Helper() + tag := MetaproxyProvisionManagedTag + baseURL := "https://old.example/v1" + priority := int64(1) + weight := uint(50) + old := Channel{ + Type: 1, + Key: "old-key", + Name: "Old upstream [old]", + BaseURL: &baseURL, + Models: "old-model", + Group: "standard", + Priority: &priority, + Weight: &weight, + Status: 2, + Tag: &tag, + UsedQuota: 12345, + CreatedTime: 99, + } + require.NoError(t, db.Create(&old).Error) + require.NoError(t, old.AddAbilities(db)) + require.NoError(t, db.Create(&[]Option{ + {Key: MetaproxyProvisionDigestOption, Value: "old-digest"}, + {Key: MetaproxyProvisionRevisionOption, Value: "old-revision"}, + {Key: "ModelRatio", Value: `{"old-model":1}`}, + {Key: "CompletionRatio", Value: `{"old-model":2}`}, + {Key: "CacheRatio", Value: `{}`}, + {Key: "GroupRatio", Value: `{"standard":1}`}, + {Key: "UserUsableGroups", Value: `{"default":"standard"}`}, + }).Error) + return old +} + +func desiredProvisionConfig() MetaproxyProvisionConfig { + return MetaproxyProvisionConfig{ + Revision: "ba57fd0526c6ce9e9869c225ac9997d96b2bbdea", + Digest: "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", + Channels: []MetaproxyProvisionChannel{ + provisionTestChannel("New upstream [new]", "new-key", "new-model"), + }, + Options: MetaproxyProvisionOptions{ + ModelRatio: `{"new-model":1.5}`, + CompletionRatio: `{"new-model":3}`, + CacheRatio: `{"new-model":0.1}`, + GroupRatio: `{"standard":1}`, + UserUsableGroups: `{"default":"standard"}`, + }, + } +} + +func TestApplyMetaproxyProvisionReplacesManagedStateAtomically(t *testing.T) { + db := setupMetaproxyProvisionTestDB(t) + seedProvisionState(t, db) + + result, err := ApplyMetaproxyProvision(desiredProvisionConfig(), "old-digest") + require.NoError(t, err) + require.False(t, result.AlreadyApplied) + require.True(t, result.RestartRequired) + require.True(t, IsMetaproxyProvisionRuntimeFrozen()) + + var channels []Channel + require.NoError(t, db.Where("tag = ?", MetaproxyProvisionManagedTag).Find(&channels).Error) + require.Len(t, channels, 1) + require.Equal(t, "New upstream [new]", channels[0].Name) + require.Equal(t, "new-key", channels[0].Key) + require.Greater(t, channels[0].CreatedTime, int64(0)) + + var abilities []Ability + require.NoError(t, db.Find(&abilities).Error) + require.Len(t, abilities, 1) + require.Equal(t, "new-model", abilities[0].Model) + require.Equal(t, channels[0].Id, abilities[0].ChannelId) + + var options []Option + require.NoError(t, db.Find(&options).Error) + got := make(map[string]string, len(options)) + for _, option := range options { + got[option.Key] = option.Value + } + require.Equal(t, desiredProvisionConfig().Digest, got[MetaproxyProvisionDigestOption]) + require.Equal(t, desiredProvisionConfig().Revision, got[MetaproxyProvisionRevisionOption]) + require.Equal(t, `{"new-model":1.5}`, got["ModelRatio"]) + + common.OptionMapRWMutex.RLock() + require.Equal(t, "old-digest", common.OptionMap[MetaproxyProvisionDigestOption]) + common.OptionMapRWMutex.RUnlock() +} + +func TestApplyMetaproxyProvisionConflictDoesNotWrite(t *testing.T) { + db := setupMetaproxyProvisionTestDB(t) + old := seedProvisionState(t, db) + + _, err := ApplyMetaproxyProvision(desiredProvisionConfig(), "some-other-digest") + require.ErrorIs(t, err, ErrMetaproxyProvisionConflict) + require.False(t, IsMetaproxyProvisionRuntimeFrozen()) + + var channel Channel + require.NoError(t, db.First(&channel, old.Id).Error) + require.Equal(t, "old-key", channel.Key) + var digest Option + require.NoError(t, db.First(&digest, "key = ?", MetaproxyProvisionDigestOption).Error) + require.Equal(t, "old-digest", digest.Value) +} + +func TestApplyMetaproxyProvisionRequiresMemoryCache(t *testing.T) { + setupMetaproxyProvisionTestDB(t) + common.MemoryCacheEnabled = false + + _, err := ApplyMetaproxyProvision(desiredProvisionConfig(), "none") + require.ErrorIs(t, err, ErrMetaproxyProvisionRequiresMemoryCache) +} + +func TestApplyMetaproxyProvisionUpdatesInPlaceAndIsIdempotent(t *testing.T) { + db := setupMetaproxyProvisionTestDB(t) + old := seedProvisionState(t, db) + config := desiredProvisionConfig() + config.Channels = []MetaproxyProvisionChannel{ + provisionTestChannel(old.Name, "rotated-key", "new-model"), + } + config.Channels[0].Status = old.Status + + first, err := ApplyMetaproxyProvision(config, "old-digest") + require.NoError(t, err) + require.False(t, first.AlreadyApplied) + + var updated Channel + require.NoError(t, db.First(&updated, old.Id).Error) + require.Equal(t, old.Id, updated.Id) + require.Equal(t, old.CreatedTime, updated.CreatedTime) + require.Equal(t, old.UsedQuota, updated.UsedQuota) + require.Equal(t, old.Status, updated.Status) + require.Equal(t, "rotated-key", updated.Key) + + second, err := ApplyMetaproxyProvision(config, "old-digest") + require.NoError(t, err) + require.True(t, second.AlreadyApplied) + require.True(t, second.RestartRequired) + + var count int64 + require.NoError(t, db.Model(&Channel{}).Where("tag = ?", MetaproxyProvisionManagedTag).Count(&count).Error) + require.EqualValues(t, 1, count) +} + +func TestApplyMetaproxyProvisionRollsBackEveryTableOnFailure(t *testing.T) { + db := setupMetaproxyProvisionTestDB(t) + old := seedProvisionState(t, db) + require.NoError(t, db.Exec(` + CREATE TRIGGER reject_completion_ratio + BEFORE UPDATE ON options + WHEN NEW.key = 'CompletionRatio' + BEGIN + SELECT RAISE(ABORT, 'injected option failure'); + END; + `).Error) + + _, err := ApplyMetaproxyProvision(desiredProvisionConfig(), "old-digest") + require.Error(t, err) + require.False(t, errors.Is(err, ErrMetaproxyProvisionConflict)) + require.False(t, IsMetaproxyProvisionRuntimeFrozen()) + + var channels []Channel + require.NoError(t, db.Where("tag = ?", MetaproxyProvisionManagedTag).Find(&channels).Error) + require.Len(t, channels, 1) + require.Equal(t, old.Id, channels[0].Id) + require.Equal(t, "old-key", channels[0].Key) + + var modelRatio Option + require.NoError(t, db.First(&modelRatio, "key = ?", "ModelRatio").Error) + require.Equal(t, `{"old-model":1}`, modelRatio.Value) + var digest Option + require.NoError(t, db.First(&digest, "key = ?", MetaproxyProvisionDigestOption).Error) + require.Equal(t, "old-digest", digest.Value) +} diff --git a/model/option.go b/model/option.go index 9b8f864d6feb..358822479184 100644 --- a/model/option.go +++ b/model/option.go @@ -226,8 +226,13 @@ func loadOptionsFromDatabase() { func SyncOptions(frequency int) { for { time.Sleep(time.Duration(frequency) * time.Second) - common.SysLog("syncing options from database") - loadOptionsFromDatabase() + if !RunMetaproxyProvisionSyncIfReady(func() { + common.SysLog("syncing options from database") + loadOptionsFromDatabase() + }) { + common.SysLog("skipping option sync while a metaproxy provision restart is pending") + continue + } } } diff --git a/router/api-router.go b/router/api-router.go index a68894dbb52f..3cf179c9760d 100644 --- a/router/api-router.go +++ b/router/api-router.go @@ -53,6 +53,13 @@ func SetApiRouter(router *gin.Engine) { // Standard OAuth providers (GitHub, Discord, OIDC, LinuxDO) - unified route apiRouter.GET("/oauth/:provider", middleware.CriticalRateLimit(), controller.HandleOAuth) apiRouter.GET("/ratio_config", middleware.CriticalRateLimit(), controller.GetRatioConfig) + apiRouter.POST( + "/metaproxy/provision", + middleware.RootAuth(), + middleware.CriticalRateLimit(), + middleware.DisableCache(), + controller.ApplyMetaproxyProvision, + ) apiRouter.POST("/stripe/webhook", anonymousRequestBodyLimit, controller.StripeWebhook) apiRouter.POST("/creem/webhook", anonymousRequestBodyLimit, controller.CreemWebhook)