feat: aiproxy modelinfo and dashboard and model rpm limit (#5291)

* feat: model info

* fix: model config vision

* feat: aiproxy dashboard api

* fix: two week and pg hour format

* fix: model tag name

* feat: model rpm limit

* fix: ci

* feat: search log with code type

* feat: resp detail buf use pool

* feat: no need init client, use ctx

* fix: lint

* feat: admin api log filed

* feat: log usage

* feat: auto retry

* fix: retry channel exhausted, use first channel

* feat: init monitor

* feat: auto ban error rate and auto test unban

* fix: getChannelWithFallback

* feat: support google thinking

* fix: monitor

* feat: get log detail

* feat: no need channel config

* feat: key validate

* feat: add model error auto ban optioon

* feat: gemini tool

* feat: gemini openai sdk

* fix: option keys

* feat: do not save access at

* fix: del no use options

* fix: del no use options

* fix: auto test banned models need return when get from redis error happend

* fix: remove channel db hook

* chore: clean detail only after insert it

* fix: err print on debug

* fix: cache update

* feat: group consume level rpm ratio

* fix: error return

* feat: decode svg

* fix: check is image

* fix: reply raw 429 message

* feat: req and resp body max size limit

* fix: _ import lint

* fix: get token encoder log

* fix: sum used amount

* fix: delete no need cache

* feat: dashboard rpm

* feat: dashboard tpm

* feat: step modelinfo

* feat: yi

* fix: yi

* feat: debug banned

* chore: bump go mod

* chore: bump go mod

* fix: save model time parse

* feat: fill dash carts gaps

* feat: fill dash carts gaps

* chore: go mod tidy

* feat: dashboard timespan

* feat: dashboard timespan from query

* feat: decouple request paths

* feat: group model tmp limit

* feat: decoupling url paths

* fix: check balance

* refactor: relay handler

* refactor: post relay

* feat: fill gaps before and after point

* fix: qwen long tokens

* feat: get rpm from redis

* fix: fill gaps

* fix: log error

* fix: token not fount err log

* fix: if err resp is not json, replay raw content

* fix: do not save same response body and content

* fix: save resp json or empty

* feat: sort distinct values

* fix: token models

* feat: redis clean expired cache

* feat: atomic model cache

* feat: consume

* feat: group custom model rpm tpm

* fix: models

* fix: v1 route

* fix: cros

* feat: rate limit err log record

* fix: rpush

* fix: dashboard time span

* feat: group model list adjusted tpm rpm

* feat: baichuan model config

* fix: rpm limit recore ignore empty channel id

* feat: disable model config

* feat: internal token

* fix: lint

* fix: recore req to redis

* feat: option from env

* fix: internal token option key

* fix: ignore redis ping error

* fix: ignore redis ping error

* fix: subscription

* fix: subscription

* feat: precheck group balance

* fix: consume nil pointer

* feat: log balance

* feat: ip log

* fix: group disable

* fix: non stream context cancel

* feat: amount log

* fix: balance and amount log format

* fix: do not skip empty

* fix: reason system prompt

* feat: doubao and moonshot model

* feat: disable model config can load existed model

* chore: add shutdown timeout duration to 600 sec

* feat: dashboard data build whit concurrent

* feat: logs data build whit concurrent

* fix: monitor remove banned model

* feat: split think

* fix: skip enpty think

* fix: do not store large resp

* fix: reat limit script

* fix: reat limit use micro second

* fix: ignore gemini input count error

* feat: calude model config

* fix: claude stream usage resp

* fix: claude stream usage resp

* fix: claude stream usage resp

* feat: auto create sqlite dir

* feat: log detail body truncated

* chore: add body conv commend

* feat: monitor ignore error rate compute when is success request

* feat: ollama usage support

* feat: baseurl embed v1 prefix

* feat: limit detail record size

* feat: split think config

* feat: channel default priority

* fix: rate limit message

* feat: channel meta api

* feat: add channel key validate help message

* fix: channel config update

* fix: split think

* fix: claude api

* fix: record total tokens

* chore: bump go mod

* chore: bump go mod

* feat: qwen open source vl models

* fix: qwen2.5 vl tool choice

* feat: stt audio duration

* feat: ali paraformer price

* fix: stt usage

* feat: qwen mt

* fix: render when split skip

* feat: sealos realname check

* feat: gemini usage support

* fix: lint

* fix: error message

* fix: lint

* fix: search token

* fix: no real name limit han message

* feat: gemini model config

* fix: get group error hans message

* fix: get group dashboard models

* feat: channel and token model search

* feat: support ali completions

* feat: internal group and search optimize

* feat: conv gemini tool choice

* fix: gemini empty tool parameters

* chore: env readme

* fix: ci lint
This commit is contained in:
zijiren
2025-02-20 10:56:32 +08:00
committed by GitHub
parent c67b4efa3f
commit 6641cd9ad3
187 changed files with 8616 additions and 5067 deletions
+14 -6
View File
@@ -1,7 +1,15 @@
FROM gcr.io/distroless/static:nonroot
ARG TARGETARCH
COPY bin/service-aiproxy-$TARGETARCH /manager
EXPOSE 3000
USER 65532:65532
FROM alpine:latest
ENTRYPOINT ["/manager"]
ARG TARGETARCH
COPY bin/service-aiproxy-$TARGETARCH /aiproxy
ENV PUID=0 PGID=0 UMASK=022
ENV FFPROBE_ENABLED=true
EXPOSE 3000
RUN apk add --no-cache ca-certificates tzdata ffmpeg && \
rm -rf /var/cache/apk/*
ENTRYPOINT ["/aiproxy"]
+10
View File
@@ -14,3 +14,13 @@ sealos run ghcr.io/labring/sealos-cloud-aiproxy-service:latest \
-e cloudDomain=<cloud-domain> \
-e LOG_SQL_DSN=""
```
# Envs
- `ADMIN_KEY`: The admin key for the AI Proxy Service, admin key is used to admin api and relay api, default is empty
- `SEALOS_JWT_KEY`: Used to sealos balance service, default is empty
- `SQL_DSN`: The database connection string, default is empty
- `LOG_SQL_DSN`: The log database connection string, default is empty
- `REDIS_CONN_STRING`: The redis connection string, default is empty
- `BALANCE_SEALOS_CHECK_REAL_NAME_ENABLE`: Whether to check real name, default is `false`
- `BALANCE_SEALOS_NO_REAL_NAME_USED_AMOUNT_LIMIT`: The amount of used balance when the user has no real name, default is `1`
+76
View File
@@ -0,0 +1,76 @@
package audio
import (
"errors"
"io"
"os/exec"
"strconv"
"strings"
"github.com/labring/sealos/service/aiproxy/common/config"
)
var ErrAudioDurationNAN = errors.New("audio duration is N/A")
func GetAudioDuration(audio io.Reader) (float64, error) {
if !config.FfprobeEnabled {
return 0, nil
}
ffprobeCmd := exec.Command(
"ffprobe",
"-v", "error",
"-select_streams", "a:0",
"-show_entries", "stream=duration",
"-of", "default=noprint_wrappers=1:nokey=1",
"-i", "-",
)
ffprobeCmd.Stdin = audio
output, err := ffprobeCmd.Output()
if err != nil {
return 0, err
}
str := strings.TrimSpace(string(output))
if str == "" || str == "N/A" {
return 0, ErrAudioDurationNAN
}
duration, err := strconv.ParseFloat(str, 64)
if err != nil {
return 0, err
}
return duration, nil
}
func GetAudioDurationFromFilePath(filePath string) (float64, error) {
if !config.FfprobeEnabled {
return 0, nil
}
ffprobeCmd := exec.Command(
"ffprobe",
"-v", "error",
"-select_streams", "a:0",
"-show_entries", "format=duration",
"-of", "default=noprint_wrappers=1:nokey=1",
"-i", filePath,
)
output, err := ffprobeCmd.Output()
if err != nil {
return 0, err
}
str := strings.TrimSpace(string(output))
if str == "" || str == "N/A" {
return 0, ErrAudioDurationNAN
}
duration, err := strconv.ParseFloat(str, 64)
if err != nil {
return 0, err
}
return duration, nil
}
+14 -4
View File
@@ -1,14 +1,24 @@
package balance
import "context"
import (
"context"
"github.com/labring/sealos/service/aiproxy/model"
)
type GroupBalance interface {
GetGroupRemainBalance(ctx context.Context, group string) (float64, PostGroupConsumer, error)
GetGroupRemainBalance(ctx context.Context, group model.GroupCache) (float64, PostGroupConsumer, error)
}
type PostGroupConsumer interface {
PostGroupConsume(ctx context.Context, tokenName string, usage float64) (float64, error)
GetBalance(ctx context.Context) (float64, error)
}
var Default GroupBalance = NewMockGroupBalance()
var (
mock GroupBalance = NewMockGroupBalance()
Default = mock
)
func MockGetGroupRemainBalance(ctx context.Context, group model.GroupCache) (float64, PostGroupConsumer, error) {
return mock.GetGroupRemainBalance(ctx, group)
}
+6 -6
View File
@@ -1,6 +1,10 @@
package balance
import "context"
import (
"context"
"github.com/labring/sealos/service/aiproxy/model"
)
var _ GroupBalance = (*MockGroupBalance)(nil)
@@ -14,14 +18,10 @@ func NewMockGroupBalance() *MockGroupBalance {
return &MockGroupBalance{}
}
func (q *MockGroupBalance) GetGroupRemainBalance(_ context.Context, _ string) (float64, PostGroupConsumer, error) {
func (q *MockGroupBalance) GetGroupRemainBalance(_ context.Context, _ model.GroupCache) (float64, PostGroupConsumer, error) {
return mockBalance, q, nil
}
func (q *MockGroupBalance) PostGroupConsume(_ context.Context, _ string, usage float64) (float64, error) {
return usage, nil
}
func (q *MockGroupBalance) GetBalance(_ context.Context) (float64, error) {
return mockBalance, nil
}
+104 -16
View File
@@ -14,6 +14,7 @@ import (
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/common/conv"
"github.com/labring/sealos/service/aiproxy/common/env"
"github.com/labring/sealos/service/aiproxy/model"
"github.com/redis/go-redis/v9"
"github.com/shopspring/decimal"
log "github.com/sirupsen/logrus"
@@ -25,6 +26,7 @@ const (
appType = "LLM-TOKEN"
sealosRequester = "sealos-admin"
sealosGroupBalanceKey = "sealos:balance:%s"
sealosUserRealNameKey = "sealos:realName:%s"
getBalanceRetry = 3
)
@@ -38,6 +40,11 @@ var (
sealosCacheExpire = 3 * time.Minute
)
var (
sealosCheckRealNameEnable = env.Bool("BALANCE_SEALOS_CHECK_REAL_NAME_ENABLE", false)
sealosNoRealNameUsedAmountLimit = env.Float64("BALANCE_SEALOS_NO_REAL_NAME_USED_AMOUNT_LIMIT", 1)
)
type Sealos struct {
accountURL string
}
@@ -145,12 +152,20 @@ func cacheDecreaseGroupBalance(ctx context.Context, group string, amount int64)
return decreaseGroupBalanceScript.Run(ctx, common.RDB, []string{fmt.Sprintf(sealosGroupBalanceKey, group)}, amount).Err()
}
func (s *Sealos) GetGroupRemainBalance(ctx context.Context, group string) (float64, PostGroupConsumer, error) {
var ErrNoRealNameUsedAmountLimit = errors.New("达到未实名用户使用额度限制,请实名认证")
func (s *Sealos) GetGroupRemainBalance(ctx context.Context, group model.GroupCache) (float64, PostGroupConsumer, error) {
var errs []error
for i := 0; ; i++ {
balance, consumer, err := s.getGroupRemainBalance(ctx, group)
balance, userUID, err := s.getGroupRemainBalance(ctx, group.ID)
if err == nil {
return balance, consumer, nil
if sealosCheckRealNameEnable &&
group.UsedAmount > sealosNoRealNameUsedAmountLimit &&
!s.checkRealName(ctx, userUID) {
return 0, nil, ErrNoRealNameUsedAmountLimit
}
return decimal.NewFromInt(balance).Div(decimalBalancePrecision).InexactFloat64(),
newSealosPostGroupConsumer(s.accountURL, group.ID, userUID), nil
}
errs = append(errs, err)
if i == getBalanceRetry-1 {
@@ -160,26 +175,105 @@ func (s *Sealos) GetGroupRemainBalance(ctx context.Context, group string) (float
}
}
func cacheGetUserRealName(ctx context.Context, userUID string) (bool, error) {
if !common.RedisEnabled || !sealosRedisCacheEnable {
return true, redis.Nil
}
ctx, cancel := context.WithTimeout(ctx, 5*time.Second)
defer cancel()
realName, err := common.RDB.Get(ctx, fmt.Sprintf(sealosUserRealNameKey, userUID)).Bool()
if err != nil {
return false, err
}
return realName, nil
}
func cacheSetUserRealName(ctx context.Context, userUID string, realName bool) error {
if !common.RedisEnabled || !sealosRedisCacheEnable {
return nil
}
ctx, cancel := context.WithTimeout(ctx, 5*time.Second)
defer cancel()
var expireTime time.Duration
if realName {
expireTime = time.Hour * 12
} else {
expireTime = time.Minute * 1
}
return common.RDB.Set(ctx, fmt.Sprintf(sealosUserRealNameKey, userUID), realName, expireTime).Err()
}
func (s *Sealos) checkRealName(ctx context.Context, userUID string) bool {
if cache, err := cacheGetUserRealName(ctx, userUID); err == nil {
return cache
} else if !errors.Is(err, redis.Nil) {
log.Errorf("get user (%s) real name cache failed: %s", userUID, err)
}
realName, err := s.fetchRealNameFromAPI(ctx, userUID)
if err != nil {
log.Errorf("fetch user (%s) real name failed: %s", userUID, err)
return true
}
if err := cacheSetUserRealName(ctx, userUID, realName); err != nil {
log.Errorf("set user (%s) real name cache failed: %s", userUID, err)
}
return realName
}
type sealosGetRealNameInfoResp struct {
IsRealName bool `json:"isRealName"`
Error string `json:"error"`
}
func (s *Sealos) fetchRealNameFromAPI(ctx context.Context, userUID string) (bool, error) {
ctx, cancel := context.WithTimeout(ctx, 5*time.Second)
defer cancel()
req, err := http.NewRequestWithContext(ctx, http.MethodGet,
fmt.Sprintf("%s/admin/v1alpha1/real-name-info?userUID=%s", s.accountURL, userUID), nil)
if err != nil {
return false, err
}
req.Header.Set("Authorization", "Bearer "+jwtToken)
resp, err := sealosHTTPClient.Do(req)
if err != nil {
return false, err
}
defer resp.Body.Close()
var sealosResp sealosGetRealNameInfoResp
if err := json.NewDecoder(resp.Body).Decode(&sealosResp); err != nil {
return false, err
}
if resp.StatusCode != http.StatusOK || sealosResp.Error != "" {
return false, fmt.Errorf("get user (%s) real name failed with status code %d, error: %s", userUID, resp.StatusCode, sealosResp.Error)
}
return sealosResp.IsRealName, nil
}
// GroupBalance interface implementation
func (s *Sealos) getGroupRemainBalance(ctx context.Context, group string) (float64, PostGroupConsumer, error) {
func (s *Sealos) getGroupRemainBalance(ctx context.Context, group string) (int64, string, error) {
if cache, err := cacheGetGroupBalance(ctx, group); err == nil && cache.UserUID != "" {
return decimal.NewFromInt(cache.Balance).Div(decimalBalancePrecision).InexactFloat64(),
newSealosPostGroupConsumer(s.accountURL, group, cache.UserUID, cache.Balance), nil
return cache.Balance, cache.UserUID, nil
} else if err != nil && !errors.Is(err, redis.Nil) {
log.Errorf("get group (%s) balance cache failed: %s", group, err)
}
balance, userUID, err := s.fetchBalanceFromAPI(ctx, group)
if err != nil {
return 0, nil, err
return 0, "", err
}
if err := cacheSetGroupBalance(ctx, group, balance, userUID); err != nil {
log.Errorf("set group (%s) balance cache failed: %s", group, err)
}
return decimal.NewFromInt(balance).Div(decimalBalancePrecision).InexactFloat64(),
newSealosPostGroupConsumer(s.accountURL, group, userUID, balance), nil
return balance, userUID, nil
}
func (s *Sealos) fetchBalanceFromAPI(ctx context.Context, group string) (balance int64, userUID string, err error) {
@@ -218,22 +312,16 @@ type SealosPostGroupConsumer struct {
accountURL string
group string
uid string
balance int64
}
func newSealosPostGroupConsumer(accountURL, group, uid string, balance int64) *SealosPostGroupConsumer {
func newSealosPostGroupConsumer(accountURL, group, uid string) *SealosPostGroupConsumer {
return &SealosPostGroupConsumer{
accountURL: accountURL,
group: group,
uid: uid,
balance: balance,
}
}
func (s *SealosPostGroupConsumer) GetBalance(_ context.Context) (float64, error) {
return decimal.NewFromInt(s.balance).Div(decimalBalancePrecision).InexactFloat64(), nil
}
func (s *SealosPostGroupConsumer) PostGroupConsume(ctx context.Context, tokenName string, usage float64) (float64, error) {
amount := s.calculateAmount(usage)
-63
View File
@@ -1,63 +0,0 @@
package client
import (
"fmt"
"net/http"
"net/url"
"time"
"github.com/labring/sealos/service/aiproxy/common/config"
log "github.com/sirupsen/logrus"
)
var (
HTTPClient *http.Client
ImpatientHTTPClient *http.Client
UserContentRequestHTTPClient *http.Client
)
func Init() {
if config.UserContentRequestProxy != "" {
log.Info(fmt.Sprintf("using %s as proxy to fetch user content", config.UserContentRequestProxy))
proxyURL, err := url.Parse(config.UserContentRequestProxy)
if err != nil {
log.Fatal("USER_CONTENT_REQUEST_PROXY set but invalid: " + config.UserContentRequestProxy)
}
transport := &http.Transport{
Proxy: http.ProxyURL(proxyURL),
}
UserContentRequestHTTPClient = &http.Client{
Transport: transport,
Timeout: time.Second * time.Duration(config.UserContentRequestTimeout),
}
} else {
UserContentRequestHTTPClient = &http.Client{}
}
var transport http.RoundTripper
if config.RelayProxy != "" {
log.Info(fmt.Sprintf("using %s as api relay proxy", config.RelayProxy))
proxyURL, err := url.Parse(config.RelayProxy)
if err != nil {
log.Fatal("USER_CONTENT_REQUEST_PROXY set but invalid: " + config.UserContentRequestProxy)
}
transport = &http.Transport{
Proxy: http.ProxyURL(proxyURL),
}
}
if config.RelayTimeout == 0 {
HTTPClient = &http.Client{
Transport: transport,
}
} else {
HTTPClient = &http.Client{
Timeout: time.Duration(config.RelayTimeout) * time.Second,
Transport: transport,
}
}
ImpatientHTTPClient = &http.Client{
Timeout: 5 * time.Second,
Transport: transport,
}
}
+106 -127
View File
@@ -1,46 +1,107 @@
package config
import (
"math"
"os"
"slices"
"strconv"
"sync"
"sync/atomic"
"time"
"github.com/labring/sealos/service/aiproxy/common/env"
)
var (
OptionMap map[string]string
OptionMapRWMutex sync.RWMutex
DebugEnabled = env.Bool("DEBUG", false)
DebugSQLEnabled = env.Bool("DEBUG_SQL", false)
)
var (
DebugEnabled, _ = strconv.ParseBool(os.Getenv("DEBUG"))
DebugSQLEnabled, _ = strconv.ParseBool(os.Getenv("DEBUG_SQL"))
DisableAutoMigrateDB = env.Bool("DISABLE_AUTO_MIGRATE_DB", false)
OnlyOneLogFile = env.Bool("ONLY_ONE_LOG_FILE", false)
AdminKey = os.Getenv("ADMIN_KEY")
FfprobeEnabled = env.Bool("FFPROBE_ENABLED", false)
)
var (
// 当测试或请求的时候发生错误是否自动禁用渠道
automaticDisableChannelEnabled atomic.Bool
// 当测试成功是否自动启用渠道
automaticEnableChannelWhenTestSucceedEnabled atomic.Bool
// 是否近似计算token
approximateTokenEnabled atomic.Bool
// 重试次数
retryTimes atomic.Int64
// 暂停服务
disableServe atomic.Bool
// log detail 存储时间(小时)
disableServe atomic.Bool
logDetailStorageHours int64 = 3 * 24
internalToken atomic.Value
)
var (
retryTimes atomic.Int64
enableModelErrorAutoBan atomic.Bool
modelErrorAutoBanRate = math.Float64bits(0.5)
timeoutWithModelType atomic.Value
disableModelConfig = env.Bool("DISABLE_MODEL_CONFIG", false)
)
var (
defaultChannelModels atomic.Value
defaultChannelModelMapping atomic.Value
groupMaxTokenNum atomic.Int64
groupConsumeLevelRatio atomic.Value
)
var geminiSafetySetting atomic.Value
var billingEnabled atomic.Bool
func init() {
timeoutWithModelType.Store(make(map[int]int64))
defaultChannelModels.Store(make(map[int][]string))
defaultChannelModelMapping.Store(make(map[int]map[string]string))
groupConsumeLevelRatio.Store(make(map[float64]float64))
geminiSafetySetting.Store("BLOCK_NONE")
billingEnabled.Store(true)
internalToken.Store(os.Getenv("INTERNAL_TOKEN"))
}
func GetDisableModelConfig() bool {
return disableModelConfig
}
func GetRetryTimes() int64 {
return retryTimes.Load()
}
func SetRetryTimes(times int64) {
times = env.Int64("RETRY_TIMES", times)
retryTimes.Store(times)
}
func GetEnableModelErrorAutoBan() bool {
return enableModelErrorAutoBan.Load()
}
func SetEnableModelErrorAutoBan(enabled bool) {
enabled = env.Bool("ENABLE_MODEL_ERROR_AUTO_BAN", enabled)
enableModelErrorAutoBan.Store(enabled)
}
func GetModelErrorAutoBanRate() float64 {
return math.Float64frombits(atomic.LoadUint64(&modelErrorAutoBanRate))
}
func SetModelErrorAutoBanRate(rate float64) {
rate = env.Float64("MODEL_ERROR_AUTO_BAN_RATE", rate)
atomic.StoreUint64(&modelErrorAutoBanRate, math.Float64bits(rate))
}
func GetTimeoutWithModelType() map[int]int64 {
return timeoutWithModelType.Load().(map[int]int64)
}
func SetTimeoutWithModelType(timeout map[int]int64) {
timeout = env.JSON("TIMEOUT_WITH_MODEL_TYPE", timeout)
timeoutWithModelType.Store(timeout)
}
func GetLogDetailStorageHours() int64 {
return atomic.LoadInt64(&logDetailStorageHours)
}
func SetLogDetailStorageHours(hours int64) {
hours = env.Int64("LOG_DETAIL_STORAGE_HOURS", hours)
atomic.StoreInt64(&logDetailStorageHours, hours)
}
@@ -49,96 +110,16 @@ func GetDisableServe() bool {
}
func SetDisableServe(disabled bool) {
disabled = env.Bool("DISABLE_SERVE", disabled)
disableServe.Store(disabled)
}
func GetAutomaticDisableChannelEnabled() bool {
return automaticDisableChannelEnabled.Load()
}
func SetAutomaticDisableChannelEnabled(enabled bool) {
automaticDisableChannelEnabled.Store(enabled)
}
func GetAutomaticEnableChannelWhenTestSucceedEnabled() bool {
return automaticEnableChannelWhenTestSucceedEnabled.Load()
}
func SetAutomaticEnableChannelWhenTestSucceedEnabled(enabled bool) {
automaticEnableChannelWhenTestSucceedEnabled.Store(enabled)
}
func GetApproximateTokenEnabled() bool {
return approximateTokenEnabled.Load()
}
func SetApproximateTokenEnabled(enabled bool) {
approximateTokenEnabled.Store(enabled)
}
func GetRetryTimes() int64 {
return retryTimes.Load()
}
func SetRetryTimes(times int64) {
retryTimes.Store(times)
}
var DisableAutoMigrateDB = os.Getenv("DISABLE_AUTO_MIGRATE_DB") == "true"
var RelayTimeout = env.Int("RELAY_TIMEOUT", 0) // unit is second
var RateLimitKeyExpirationDuration = 20 * time.Minute
var OnlyOneLogFile = env.Bool("ONLY_ONE_LOG_FILE", false)
var (
// 代理地址
RelayProxy = env.String("RELAY_PROXY", "")
// 用户内容请求代理地址
UserContentRequestProxy = env.String("USER_CONTENT_REQUEST_PROXY", "")
// 用户内容请求超时时间,单位为秒
UserContentRequestTimeout = env.Int("USER_CONTENT_REQUEST_TIMEOUT", 30)
)
var AdminKey = env.String("ADMIN_KEY", "")
var (
globalAPIRateLimitNum atomic.Int64
defaultChannelModels atomic.Value
defaultChannelModelMapping atomic.Value
defaultGroupQPM atomic.Int64
groupMaxTokenNum atomic.Int32
)
func init() {
defaultChannelModels.Store(make(map[int][]string))
defaultChannelModelMapping.Store(make(map[int]map[string]string))
}
// 全局qpm,不是根据ip限制,而是所有请求共享一个qpm
func GetGlobalAPIRateLimitNum() int64 {
return globalAPIRateLimitNum.Load()
}
func SetGlobalAPIRateLimitNum(num int64) {
globalAPIRateLimitNum.Store(num)
}
// group默认qpm,如果group没有设置qpm,则使用该qpm
func GetDefaultGroupQPM() int64 {
return defaultGroupQPM.Load()
}
func SetDefaultGroupQPM(qpm int64) {
defaultGroupQPM.Store(qpm)
}
func GetDefaultChannelModels() map[int][]string {
return defaultChannelModels.Load().(map[int][]string)
}
func SetDefaultChannelModels(models map[int][]string) {
models = env.JSON("DEFAULT_CHANNEL_MODELS", models)
for key, ms := range models {
slices.Sort(ms)
models[key] = slices.Compact(ms)
@@ -151,54 +132,52 @@ func GetDefaultChannelModelMapping() map[int]map[string]string {
}
func SetDefaultChannelModelMapping(mapping map[int]map[string]string) {
mapping = env.JSON("DEFAULT_CHANNEL_MODEL_MAPPING", mapping)
defaultChannelModelMapping.Store(mapping)
}
// 那个group最多可创建的token数量,0表示不限制
func GetGroupMaxTokenNum() int32 {
func GetGroupConsumeLevelRatio() map[float64]float64 {
return groupConsumeLevelRatio.Load().(map[float64]float64)
}
func SetGroupConsumeLevelRatio(ratio map[float64]float64) {
ratio = env.JSON("GROUP_CONSUME_LEVEL_RATIO", ratio)
groupConsumeLevelRatio.Store(ratio)
}
// GetGroupMaxTokenNum returns max number of tokens per group, 0 means unlimited
func GetGroupMaxTokenNum() int64 {
return groupMaxTokenNum.Load()
}
func SetGroupMaxTokenNum(num int32) {
func SetGroupMaxTokenNum(num int64) {
num = env.Int64("GROUP_MAX_TOKEN_NUM", num)
groupMaxTokenNum.Store(num)
}
var (
geminiSafetySetting atomic.Value
geminiVersion atomic.Value
)
func init() {
geminiSafetySetting.Store("BLOCK_NONE")
geminiVersion.Store("v1beta")
}
func GetGeminiSafetySetting() string {
return geminiSafetySetting.Load().(string)
}
func SetGeminiSafetySetting(setting string) {
setting = env.String("GEMINI_SAFETY_SETTING", setting)
geminiSafetySetting.Store(setting)
}
func GetGeminiVersion() string {
return geminiVersion.Load().(string)
}
func SetGeminiVersion(version string) {
geminiVersion.Store(version)
}
var billingEnabled atomic.Bool
func init() {
billingEnabled.Store(true)
}
func GetBillingEnabled() bool {
return billingEnabled.Load()
}
func SetBillingEnabled(enabled bool) {
enabled = env.Bool("BILLING_ENABLED", enabled)
billingEnabled.Store(enabled)
}
func GetInternalToken() string {
return internalToken.Load().(string)
}
func SetInternalToken(token string) {
token = env.String("INTERNAL_TOKEN", token)
internalToken.Store(token)
}
+181
View File
@@ -0,0 +1,181 @@
package consume
import (
"context"
"sync"
"github.com/labring/sealos/service/aiproxy/common/balance"
"github.com/labring/sealos/service/aiproxy/model"
"github.com/labring/sealos/service/aiproxy/relay/meta"
relaymodel "github.com/labring/sealos/service/aiproxy/relay/model"
"github.com/shopspring/decimal"
log "github.com/sirupsen/logrus"
)
var consumeWaitGroup sync.WaitGroup
func Wait() {
consumeWaitGroup.Wait()
}
func AsyncConsume(
postGroupConsumer balance.PostGroupConsumer,
code int,
usage *relaymodel.Usage,
meta *meta.Meta,
inputPrice,
outputPrice float64,
content string,
ip string,
requestDetail *model.RequestDetail,
) {
if meta.IsChannelTest {
return
}
consumeWaitGroup.Add(1)
defer func() {
consumeWaitGroup.Done()
if r := recover(); r != nil {
log.Errorf("panic in consume: %v", r)
}
}()
go Consume(
context.Background(),
postGroupConsumer,
code,
usage,
meta,
inputPrice,
outputPrice,
content,
ip,
requestDetail,
)
}
func Consume(
ctx context.Context,
postGroupConsumer balance.PostGroupConsumer,
code int,
usage *relaymodel.Usage,
meta *meta.Meta,
inputPrice,
outputPrice float64,
content string,
ip string,
requestDetail *model.RequestDetail,
) {
if meta.IsChannelTest {
return
}
amount := CalculateAmount(usage, inputPrice, outputPrice)
amount = consumeAmount(ctx, amount, postGroupConsumer, meta)
err := recordConsume(meta, code, usage, inputPrice, outputPrice, content, ip, requestDetail, amount)
if err != nil {
log.Error("error batch record consume: " + err.Error())
}
}
func consumeAmount(
ctx context.Context,
amount float64,
postGroupConsumer balance.PostGroupConsumer,
meta *meta.Meta,
) float64 {
if amount > 0 {
return processGroupConsume(ctx, amount, postGroupConsumer, meta)
}
return 0
}
func CalculateAmount(
usage *relaymodel.Usage,
inputPrice, outputPrice float64,
) float64 {
if usage == nil {
return 0
}
promptTokens := usage.PromptTokens
completionTokens := usage.CompletionTokens
totalTokens := promptTokens + completionTokens
if totalTokens == 0 {
return 0
}
promptAmount := decimal.NewFromInt(int64(promptTokens)).
Mul(decimal.NewFromFloat(inputPrice)).
Div(decimal.NewFromInt(model.PriceUnit))
completionAmount := decimal.NewFromInt(int64(completionTokens)).
Mul(decimal.NewFromFloat(outputPrice)).
Div(decimal.NewFromInt(model.PriceUnit))
return promptAmount.Add(completionAmount).InexactFloat64()
}
func processGroupConsume(
ctx context.Context,
amount float64,
postGroupConsumer balance.PostGroupConsumer,
meta *meta.Meta,
) float64 {
consumedAmount, err := postGroupConsumer.PostGroupConsume(ctx, meta.Token.Name, amount)
if err != nil {
log.Error("error consuming token remain amount: " + err.Error())
if err := model.CreateConsumeError(
meta.RequestID,
meta.RequestAt,
meta.Group.ID,
meta.Token.Name,
meta.OriginModel,
err.Error(),
amount,
meta.Token.ID,
); err != nil {
log.Error("failed to create consume error: " + err.Error())
}
return amount
}
return consumedAmount
}
func recordConsume(meta *meta.Meta, code int, usage *relaymodel.Usage, inputPrice, outputPrice float64, content string, ip string, requestDetail *model.RequestDetail, amount float64) error {
promptTokens := 0
completionTokens := 0
if usage != nil {
promptTokens = usage.PromptTokens
completionTokens = usage.CompletionTokens
}
var channelID int
if meta.Channel != nil {
channelID = meta.Channel.ID
}
return model.BatchRecordConsume(
meta.RequestID,
meta.RequestAt,
meta.Group.ID,
code,
channelID,
promptTokens,
completionTokens,
meta.OriginModel,
meta.Token.ID,
meta.Token.Name,
amount,
inputPrice,
outputPrice,
meta.Endpoint,
content,
meta.Mode,
ip,
requestDetail,
)
}
+2 -7
View File
@@ -9,15 +9,10 @@ func AsString(v any) string {
// The change of bytes will cause the change of string synchronously
func BytesToString(b []byte) string {
return *(*string)(unsafe.Pointer(&b))
return unsafe.String(unsafe.SliceData(b), len(b))
}
// If string is readonly, modifying bytes will cause panic
func StringToBytes(s string) []byte {
return *(*[]byte)(unsafe.Pointer(
&struct {
string
Cap int
}{s, len(s)},
))
return unsafe.Slice(unsafe.StringData(s), len(s))
}
+7 -9
View File
@@ -1,13 +1,11 @@
package ctxkey
type OriginalModelKey string
const (
OriginalModel OriginalModelKey = "original_model"
)
const (
Channel = "channel"
Group = "group"
Token = "token"
Group = "group"
Token = "token"
GroupBalance = "group_balance"
OriginalModel = "original_model"
RequestID = "X-Request-Id"
ModelCaches = "model_caches"
ModelConfig = "model_config"
)
+2 -2
View File
@@ -11,6 +11,6 @@ var (
)
var (
SQLitePath = "aiproxy.db"
SQLiteBusyTimeout = env.Int("SQLITE_BUSY_TIMEOUT", 3000)
SQLitePath = env.String("SQLITE_PATH", "aiproxy.db")
SQLiteBusyTimeout = env.Int64("SQLITE_BUSY_TIMEOUT", 3000)
)
+52 -9
View File
@@ -3,40 +3,83 @@ package env
import (
"os"
"strconv"
json "github.com/json-iterator/go"
"github.com/labring/sealos/service/aiproxy/common/conv"
log "github.com/sirupsen/logrus"
)
func Bool(env string, defaultValue bool) bool {
if env == "" || os.Getenv(env) == "" {
if env == "" {
return defaultValue
}
return os.Getenv(env) == "true"
e := os.Getenv(env)
if e == "" {
return defaultValue
}
p, err := strconv.ParseBool(e)
if err != nil {
log.Errorf("invalid %s: %s", env, e)
return defaultValue
}
return p
}
func Int(env string, defaultValue int) int {
if env == "" || os.Getenv(env) == "" {
func Int64(env string, defaultValue int64) int64 {
if env == "" {
return defaultValue
}
num, err := strconv.Atoi(os.Getenv(env))
e := os.Getenv(env)
if e == "" {
return defaultValue
}
num, err := strconv.ParseInt(e, 10, 64)
if err != nil {
log.Errorf("invalid %s: %s", env, e)
return defaultValue
}
return num
}
func Float64(env string, defaultValue float64) float64 {
if env == "" || os.Getenv(env) == "" {
if env == "" {
return defaultValue
}
num, err := strconv.ParseFloat(os.Getenv(env), 64)
e := os.Getenv(env)
if e == "" {
return defaultValue
}
num, err := strconv.ParseFloat(e, 64)
if err != nil {
log.Errorf("invalid %s: %s", env, e)
return defaultValue
}
return num
}
func String(env string, defaultValue string) string {
if env == "" || os.Getenv(env) == "" {
if env == "" {
return defaultValue
}
return os.Getenv(env)
e := os.Getenv(env)
if e == "" {
return defaultValue
}
return e
}
func JSON[T any](env string, defaultValue T) T {
if env == "" {
return defaultValue
}
e := os.Getenv(env)
if e == "" {
return defaultValue
}
var t T
if err := json.Unmarshal(conv.StringToBytes(e), &t); err != nil {
log.Errorf("invalid %s: %s", env, e)
return defaultValue
}
return t
}
@@ -7,7 +7,6 @@ import (
json "github.com/json-iterator/go"
"github.com/labring/sealos/service/aiproxy/common/conv"
"gorm.io/gorm/schema"
)
+43 -1
View File
@@ -3,9 +3,11 @@ package common
import (
"bytes"
"context"
"errors"
"fmt"
"io"
"net/http"
"strings"
"github.com/gin-gonic/gin"
json "github.com/json-iterator/go"
@@ -13,7 +15,38 @@ import (
type RequestBodyKey struct{}
const (
MaxRequestBodySize = 1024 * 1024 * 50 // 50MB
)
func LimitReader(r io.Reader, n int64) io.Reader { return &LimitedReader{r, n} }
type LimitedReader struct {
R io.Reader
N int64
}
var ErrLimitedReaderExceeded = errors.New("limited reader exceeded")
func (l *LimitedReader) Read(p []byte) (n int, err error) {
if l.N <= 0 {
return 0, ErrLimitedReaderExceeded
}
if int64(len(p)) > l.N {
p = p[0:l.N]
}
n, err = l.R.Read(p)
l.N -= int64(n)
return
}
func GetRequestBody(req *http.Request) ([]byte, error) {
contentType := req.Header.Get("Content-Type")
if contentType == "application/x-www-form-urlencoded" ||
strings.HasPrefix(contentType, "multipart/form-data") {
return nil, nil
}
requestBody := req.Context().Value(RequestBodyKey{})
if requestBody != nil {
return requestBody.([]byte), nil
@@ -27,8 +60,17 @@ func GetRequestBody(req *http.Request) ([]byte, error) {
}
}()
if req.ContentLength <= 0 || req.Header.Get("Content-Type") != "application/json" {
buf, err = io.ReadAll(req.Body)
buf, err = io.ReadAll(LimitReader(req.Body, MaxRequestBodySize))
if err != nil {
if errors.Is(err, ErrLimitedReaderExceeded) {
return nil, fmt.Errorf("request body too large, max: %d", MaxRequestBodySize)
}
return nil, fmt.Errorf("request body read failed: %w", err)
}
} else {
if req.ContentLength > MaxRequestBodySize {
return nil, fmt.Errorf("request body too large: %d, max: %d", req.ContentLength, MaxRequestBodySize)
}
buf = make([]byte, req.ContentLength)
_, err = io.ReadFull(req.Body, buf)
}
-41
View File
@@ -1,41 +0,0 @@
package helper
import (
"fmt"
"strconv"
"time"
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/common/random"
)
func GenRequestID() string {
return strconv.FormatInt(time.Now().UnixMilli(), 10) + random.GetRandomNumberString(4)
}
func GetResponseID(c *gin.Context) string {
logID := c.GetString(string(RequestIDKey))
return "chatcmpl-" + logID
}
func AssignOrDefault(value string, defaultValue string) string {
if len(value) != 0 {
return value
}
return defaultValue
}
func MessageWithRequestID(message string, id string) string {
return fmt.Sprintf("%s (request id: %s)", message, id)
}
func String2Int(keyword string) int {
if keyword == "" {
return 0
}
i, err := strconv.Atoi(keyword)
if err != nil {
return 0
}
return i
}
-7
View File
@@ -1,7 +0,0 @@
package helper
type Key string
const (
RequestIDKey Key = "X-Request-Id"
)
-9
View File
@@ -1,9 +0,0 @@
package helper
import (
"time"
)
func GetTimestamp() int64 {
return time.Now().Unix()
}
+22 -9
View File
@@ -19,10 +19,10 @@ import (
"regexp"
"strings"
"github.com/labring/sealos/service/aiproxy/common"
// import webp decoder
_ "golang.org/x/image/webp"
"github.com/labring/sealos/service/aiproxy/common/client"
)
// Regex to match data URL pattern
@@ -37,7 +37,7 @@ func GetImageSizeFromURL(url string) (width int, height int, err error) {
if err != nil {
return 0, 0, err
}
resp, err := client.UserContentRequestHTTPClient.Do(req)
resp, err := http.DefaultClient.Do(req)
if err != nil {
return
}
@@ -58,6 +58,10 @@ func GetImageSizeFromURL(url string) (width int, height int, err error) {
return img.Width, img.Height, nil
}
const (
MaxImageSize = 1024 * 1024 * 5 // 5MB
)
func GetImageFromURL(ctx context.Context, url string) (string, string, error) {
// Check if the URL is a data URL
matches := dataURLPattern.FindStringSubmatch(url)
@@ -70,7 +74,7 @@ func GetImageFromURL(ctx context.Context, url string) (string, string, error) {
if err != nil {
return "", "", err
}
resp, err := client.UserContentRequestHTTPClient.Do(req)
resp, err := http.DefaultClient.Do(req)
if err != nil {
return "", "", err
}
@@ -78,20 +82,29 @@ func GetImageFromURL(ctx context.Context, url string) (string, string, error) {
if resp.StatusCode != http.StatusOK {
return "", "", fmt.Errorf("status code: %d", resp.StatusCode)
}
isImage := IsImageURL(resp)
if !isImage {
return "", "", errors.New("not an image")
}
var buf []byte
if resp.ContentLength <= 0 {
buf, err = io.ReadAll(resp.Body)
buf, err = io.ReadAll(common.LimitReader(resp.Body, MaxImageSize))
if err != nil {
if errors.Is(err, common.ErrLimitedReaderExceeded) {
return "", "", fmt.Errorf("image too large, max: %d", MaxImageSize)
}
return "", "", fmt.Errorf("image read failed: %w", err)
}
} else {
if resp.ContentLength > MaxImageSize {
return "", "", fmt.Errorf("image too large: %d, max: %d", resp.ContentLength, MaxImageSize)
}
buf = make([]byte, resp.ContentLength)
_, err = io.ReadFull(resp.Body, buf)
}
if err != nil {
return "", "", err
}
isImage := IsImageURL(resp)
if !isImage {
return "", "", errors.New("not an image")
}
return resp.Header.Get("Content-Type"), base64.StdEncoding.EncodeToString(buf), nil
}
-176
View File
@@ -1,176 +0,0 @@
package image_test
import (
"encoding/base64"
"image"
_ "image/gif"
_ "image/jpeg"
_ "image/png"
"io"
"net/http"
"strconv"
"strings"
"testing"
"github.com/labring/sealos/service/aiproxy/common/client"
img "github.com/labring/sealos/service/aiproxy/common/image"
"github.com/stretchr/testify/assert"
_ "golang.org/x/image/webp"
)
type CountingReader struct {
reader io.Reader
BytesRead int
}
func (r *CountingReader) Read(p []byte) (n int, err error) {
n, err = r.reader.Read(p)
r.BytesRead += n
return n, err
}
var cases = []struct {
url string
format string
width int
height int
}{
{"https://upload.wikimedia.org/wikipedia/commons/thumb/d/dd/Gfp-wisconsin-madison-the-nature-boardwalk.jpg/2560px-Gfp-wisconsin-madison-the-nature-boardwalk.jpg", "jpeg", 2560, 1669},
{"https://upload.wikimedia.org/wikipedia/commons/9/97/Basshunter_live_performances.png", "png", 4500, 2592},
{"https://upload.wikimedia.org/wikipedia/commons/c/c6/TO_THE_ONE_SOMETHINGNESS.webp", "webp", 984, 985},
{"https://upload.wikimedia.org/wikipedia/commons/d/d0/01_Das_Sandberg-Modell.gif", "gif", 1917, 1533},
{"https://upload.wikimedia.org/wikipedia/commons/6/62/102Cervus.jpg", "jpeg", 270, 230},
}
func TestMain(m *testing.M) {
client.Init()
m.Run()
}
func TestDecode(t *testing.T) {
// Bytes read: varies sometimes
// jpeg: 1063892
// png: 294462
// webp: 99529
// gif: 956153
// jpeg#01: 32805
for _, c := range cases {
t.Run("Decode:"+c.format, func(t *testing.T) {
resp, err := http.Get(c.url)
assert.NoError(t, err)
defer resp.Body.Close()
reader := &CountingReader{reader: resp.Body}
img, format, err := image.Decode(reader)
assert.NoError(t, err)
size := img.Bounds().Size()
assert.Equal(t, c.format, format)
assert.Equal(t, c.width, size.X)
assert.Equal(t, c.height, size.Y)
t.Logf("Bytes read: %d", reader.BytesRead)
})
}
// Bytes read:
// jpeg: 4096
// png: 4096
// webp: 4096
// gif: 4096
// jpeg#01: 4096
for _, c := range cases {
t.Run("DecodeConfig:"+c.format, func(t *testing.T) {
resp, err := http.Get(c.url)
assert.NoError(t, err)
defer resp.Body.Close()
reader := &CountingReader{reader: resp.Body}
config, format, err := image.DecodeConfig(reader)
assert.NoError(t, err)
assert.Equal(t, c.format, format)
assert.Equal(t, c.width, config.Width)
assert.Equal(t, c.height, config.Height)
t.Logf("Bytes read: %d", reader.BytesRead)
})
}
}
func TestBase64(t *testing.T) {
// Bytes read:
// jpeg: 1063892
// png: 294462
// webp: 99072
// gif: 953856
// jpeg#01: 32805
for _, c := range cases {
t.Run("Decode:"+c.format, func(t *testing.T) {
resp, err := http.Get(c.url)
assert.NoError(t, err)
defer resp.Body.Close()
data, err := io.ReadAll(resp.Body)
assert.NoError(t, err)
encoded := base64.StdEncoding.EncodeToString(data)
body := base64.NewDecoder(base64.StdEncoding, strings.NewReader(encoded))
reader := &CountingReader{reader: body}
img, format, err := image.Decode(reader)
assert.NoError(t, err)
size := img.Bounds().Size()
assert.Equal(t, c.format, format)
assert.Equal(t, c.width, size.X)
assert.Equal(t, c.height, size.Y)
t.Logf("Bytes read: %d", reader.BytesRead)
})
}
// Bytes read:
// jpeg: 1536
// png: 768
// webp: 768
// gif: 1536
// jpeg#01: 3840
for _, c := range cases {
t.Run("DecodeConfig:"+c.format, func(t *testing.T) {
resp, err := http.Get(c.url)
assert.NoError(t, err)
defer resp.Body.Close()
data, err := io.ReadAll(resp.Body)
assert.NoError(t, err)
encoded := base64.StdEncoding.EncodeToString(data)
body := base64.NewDecoder(base64.StdEncoding, strings.NewReader(encoded))
reader := &CountingReader{reader: body}
config, format, err := image.DecodeConfig(reader)
assert.NoError(t, err)
assert.Equal(t, c.format, format)
assert.Equal(t, c.width, config.Width)
assert.Equal(t, c.height, config.Height)
t.Logf("Bytes read: %d", reader.BytesRead)
})
}
}
func TestGetImageSize(t *testing.T) {
for i, c := range cases {
t.Run("Decode:"+strconv.Itoa(i), func(t *testing.T) {
width, height, err := img.GetImageSize(c.url)
assert.NoError(t, err)
assert.Equal(t, c.width, width)
assert.Equal(t, c.height, height)
})
}
}
func TestGetImageSizeFromBase64(t *testing.T) {
for i, c := range cases {
t.Run("Decode:"+strconv.Itoa(i), func(t *testing.T) {
resp, err := http.Get(c.url)
assert.NoError(t, err)
defer resp.Body.Close()
data, err := io.ReadAll(resp.Body)
assert.NoError(t, err)
encoded := base64.StdEncoding.EncodeToString(data)
width, height, err := img.GetImageSizeFromBase64(encoded)
assert.NoError(t, err)
assert.Equal(t, c.width, width)
assert.Equal(t, c.height, height)
})
}
}
+45
View File
@@ -0,0 +1,45 @@
package image
import (
"image"
"image/color"
"io"
"github.com/srwiley/oksvg"
"github.com/srwiley/rasterx"
)
func Decode(r io.Reader) (image.Image, error) {
icon, err := oksvg.ReadIconStream(r)
if err != nil {
return nil, err
}
w, h := int(icon.ViewBox.W), int(icon.ViewBox.H)
icon.SetTarget(0, 0, float64(w), float64(h))
rgba := image.NewRGBA(image.Rect(0, 0, w, h))
icon.Draw(rasterx.NewDasher(w, h, rasterx.NewScannerGV(w, h, rgba, rgba.Bounds())), 1)
return rgba, err
}
func DecodeConfig(r io.Reader) (image.Config, error) {
var config image.Config
icon, err := oksvg.ReadIconStream(r)
if err != nil {
return config, err
}
config.ColorModel = color.RGBAModel
config.Width = int(icon.ViewBox.W)
config.Height = int(icon.ViewBox.H)
return config, nil
}
func init() {
image.RegisterFormat("svg", "<?xml ", Decode, DecodeConfig)
image.RegisterFormat("svg", "<svg", Decode, DecodeConfig)
}
-3
View File
@@ -16,9 +16,6 @@ var (
func Init() {
flag.Parse()
if os.Getenv("SQLITE_PATH") != "" {
SQLitePath = os.Getenv("SQLITE_PATH")
}
if *LogDir != "" {
var err error
*LogDir, err = filepath.Abs(*LogDir)
+5 -1
View File
@@ -32,7 +32,11 @@ func InitRedisClient() (err error) {
defer cancel()
_, err = RDB.Ping(ctx).Result()
return err
if err != nil {
log.Errorf("failed to ping redis: %s", err.Error())
}
return nil
}
func RedisSet(key string, value string, expiration time.Duration) error {
@@ -1,4 +1,4 @@
package common
package rpmlimit
import (
"sync"
@@ -0,0 +1,132 @@
package rpmlimit
import (
"context"
"fmt"
"time"
"github.com/labring/sealos/service/aiproxy/common"
log "github.com/sirupsen/logrus"
)
var inMemoryRateLimiter InMemoryRateLimiter
const (
groupModelRPMKey = "group_model_rpm:%s:%s"
)
var pushRequestScript = `
local key = KEYS[1]
local window = tonumber(ARGV[1])
local current_time = tonumber(ARGV[2])
local cutoff = current_time - window
redis.call('ZREMRANGEBYSCORE', key, '-inf', cutoff)
redis.call('ZADD', key, current_time, current_time)
redis.call('PEXPIRE', key, window)
return redis.call('ZCOUNT', key, cutoff, current_time)
`
var getRequestCountScript = `
local pattern = ARGV[1]
local window = tonumber(ARGV[2])
local current_time = tonumber(ARGV[3])
local cutoff = current_time - window
local keys = redis.call('KEYS', pattern)
local total = 0
for _, key in ipairs(keys) do
redis.call('ZREMRANGEBYSCORE', key, '-inf', cutoff)
local count = redis.call('ZCOUNT', key, cutoff, current_time)
total = total + count
end
return total
`
func GetRPM(ctx context.Context, group, model string) (int64, error) {
if !common.RedisEnabled {
return 0, nil
}
var pattern string
if group == "" && model == "" {
pattern = "group_model_rpm:*:*"
} else if group == "" {
pattern = "group_model_rpm:*:" + model
} else if model == "" {
pattern = fmt.Sprintf("group_model_rpm:%s:*", group)
} else {
pattern = fmt.Sprintf("group_model_rpm:%s:%s", group, model)
}
rdb := common.RDB
result, err := rdb.Eval(
ctx,
getRequestCountScript,
[]string{},
pattern,
time.Minute.Microseconds(),
time.Now().UnixMicro(),
).Int64()
if err != nil {
return 0, err
}
return result, nil
}
func redisRateLimitRequest(ctx context.Context, group, model string, maxRequestNum int64, duration time.Duration) (bool, error) {
result, err := PushRequest(ctx, group, model, duration)
if err != nil {
return false, err
}
return result <= maxRequestNum, nil
}
func PushRequest(ctx context.Context, group, model string, duration time.Duration) (int64, error) {
result, err := common.RDB.Eval(
ctx,
pushRequestScript,
[]string{
fmt.Sprintf(groupModelRPMKey, group, model),
},
duration.Microseconds(),
time.Now().UnixMicro(),
).Int64()
if err != nil {
return 0, err
}
return result, nil
}
func RateLimit(ctx context.Context, group, model string, maxRequestNum int64, duration time.Duration) (bool, error) {
if maxRequestNum == 0 {
return true, nil
}
if common.RedisEnabled {
return redisRateLimitRequest(ctx, group, model, maxRequestNum, duration)
}
return MemoryRateLimit(ctx, group, model, maxRequestNum, duration), nil
}
// ignore redis error
func ForceRateLimit(ctx context.Context, group, model string, maxRequestNum int64, duration time.Duration) bool {
if maxRequestNum == 0 {
return true
}
if common.RedisEnabled {
ok, err := redisRateLimitRequest(ctx, group, model, maxRequestNum, duration)
if err == nil {
return ok
}
log.Error("rate limit error: " + err.Error())
}
return MemoryRateLimit(ctx, group, model, maxRequestNum, duration)
}
func MemoryRateLimit(_ context.Context, group, model string, maxRequestNum int64, duration time.Duration) bool {
// It's safe to call multi times.
inMemoryRateLimiter.Init(3 * time.Minute)
return inMemoryRateLimiter.Request(fmt.Sprintf(groupModelRPMKey, group, model), int(maxRequestNum), duration)
}
+122
View File
@@ -0,0 +1,122 @@
package splitter
import "bytes"
type Splitter struct {
head []byte
tail []byte
headLen int
tailLen int
buffer []byte
state int
partialTailPos int
kmpNext []int
}
func NewSplitter(head, tail []byte) *Splitter {
return &Splitter{
head: head,
tail: tail,
headLen: len(head),
tailLen: len(tail),
kmpNext: computeKMPNext(tail),
}
}
func computeKMPNext(pattern []byte) []int {
n := len(pattern)
next := make([]int, n)
if n == 0 {
return next
}
next[0] = 0
for i := 1; i < n; i++ {
j := next[i-1]
for j > 0 && pattern[i] != pattern[j] {
j = next[j-1]
}
if pattern[i] == pattern[j] {
j++
}
next[i] = j
}
return next
}
func (s *Splitter) Process(data []byte) ([]byte, []byte) {
if len(data) == 0 {
return nil, nil
}
switch s.state {
case 0:
s.buffer = append(s.buffer, data...)
bufferLen := len(s.buffer)
minLen := bufferLen
if minLen > s.headLen {
minLen = s.headLen
}
if minLen > 0 {
if !bytes.Equal(s.buffer[:minLen], s.head[:minLen]) {
s.state = 2
remaining := s.buffer
s.buffer = nil
return nil, remaining
}
}
if bufferLen < s.headLen {
return nil, nil
}
s.state = 1
s.buffer = s.buffer[s.headLen:]
if len(s.buffer) == 0 {
return nil, nil
}
return s.processSeekTail()
case 1:
s.buffer = append(s.buffer, data...)
return s.processSeekTail()
default:
return nil, data
}
}
func (s *Splitter) processSeekTail() ([]byte, []byte) {
data := s.buffer
j := s.partialTailPos
tail := s.tail
tailLen := s.tailLen
kmpNext := s.kmpNext
var i int
for i = 0; i < len(data); i++ {
for j > 0 && data[i] != tail[j] {
j = kmpNext[j-1]
}
if data[i] == tail[j] {
j++
if j == tailLen {
end := i - tailLen + 1
if end < 0 {
end = 0
}
result := data[:end]
remaining := data[i+1:]
s.buffer = nil
s.state = 2
s.partialTailPos = 0
return result, remaining
}
}
}
splitAt := len(data) - j
if splitAt < 0 {
splitAt = 0
}
result := data[:splitAt]
remainingPart := data[splitAt:]
s.partialTailPos = j
s.buffer = remainingPart
return result, nil
}
+17
View File
@@ -0,0 +1,17 @@
package splitter
import "github.com/labring/sealos/service/aiproxy/common/conv"
const (
ThinkHead = "<think>\n"
ThinkTail = "</think>\n"
)
var (
thinkHeadBytes = conv.StringToBytes(ThinkHead)
thinkTailBytes = conv.StringToBytes(ThinkTail)
)
func NewThinkSplitter() *Splitter {
return NewSplitter(thinkHeadBytes, thinkTailBytes)
}
+31
View File
@@ -0,0 +1,31 @@
package common
import (
"unicode/utf8"
"github.com/labring/sealos/service/aiproxy/common/conv"
)
func TruncateByRune(s string, length int) string {
total := 0
for _, r := range s {
runeLen := utf8.RuneLen(r)
if runeLen == -1 || total+runeLen > length {
return s[:total]
}
total += runeLen
}
return s[:total]
}
func TruncateBytesByRune(b []byte, length int) []byte {
total := 0
for _, r := range conv.BytesToString(b) {
runeLen := utf8.RuneLen(r)
if runeLen == -1 || total+runeLen > length {
return b[:total]
}
total += runeLen
}
return b[:total]
}
+18 -18
View File
@@ -1,21 +1,20 @@
package controller
import (
"errors"
"fmt"
"net/http"
"strconv"
"time"
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/common/balance"
"github.com/labring/sealos/service/aiproxy/common/ctxkey"
"github.com/labring/sealos/service/aiproxy/middleware"
"github.com/labring/sealos/service/aiproxy/model"
"github.com/labring/sealos/service/aiproxy/relay/adaptor"
"github.com/labring/sealos/service/aiproxy/relay/adaptor/openai"
"github.com/labring/sealos/service/aiproxy/relay/channeltype"
log "github.com/sirupsen/logrus"
"github.com/gin-gonic/gin"
)
// https://github.com/labring/sealos/service/aiproxy/issues/79
@@ -25,7 +24,7 @@ func updateChannelBalance(channel *model.Channel) (float64, error) {
if !ok {
return 0, fmt.Errorf("invalid channel type: %d", channel.Type)
}
if getBalance, ok := adaptorI.(adaptor.GetBalance); ok {
if getBalance, ok := adaptorI.(adaptor.Balancer); ok {
balance, err := getBalance.GetBalance(channel)
if err != nil {
return 0, err
@@ -48,7 +47,7 @@ func UpdateChannelBalance(c *gin.Context) {
})
return
}
channel, err := model.GetChannelByID(id, false)
channel, err := model.GetChannelByID(id)
if err != nil {
c.JSON(http.StatusOK, middleware.APIResponse{
Success: false,
@@ -72,7 +71,7 @@ func UpdateChannelBalance(c *gin.Context) {
}
func updateAllChannelsBalance() error {
channels, err := model.GetAllChannels(false, false)
channels, err := model.GetAllChannels()
if err != nil {
return err
}
@@ -105,29 +104,30 @@ func AutomaticallyUpdateChannels(frequency int) {
// subscription
func GetSubscription(c *gin.Context) {
group := c.MustGet(ctxkey.Group).(*model.GroupCache)
b, _, err := balance.Default.GetGroupRemainBalance(c, group.ID)
group := middleware.GetGroup(c)
b, _, err := balance.Default.GetGroupRemainBalance(c, *group)
if err != nil {
if errors.Is(err, balance.ErrNoRealNameUsedAmountLimit) {
middleware.ErrorResponse(c, http.StatusForbidden, err.Error())
return
}
log.Errorf("get group (%s) balance failed: %s", group.ID, err)
c.JSON(http.StatusOK, middleware.APIResponse{
Success: false,
Message: fmt.Sprintf("get group (%s) balance failed", group.ID),
})
middleware.ErrorResponse(c, http.StatusInternalServerError, fmt.Sprintf("get group (%s) balance failed", group.ID))
return
}
token := c.MustGet(ctxkey.Token).(*model.TokenCache)
token := middleware.GetToken(c)
quota := token.Quota
if quota <= 0 {
quota = b
}
c.JSON(http.StatusOK, openai.SubscriptionResponse{
HardLimitUSD: quota / 7,
SoftLimitUSD: b / 7,
SystemHardLimitUSD: quota / 7,
HardLimitUSD: quota + token.UsedAmount,
SoftLimitUSD: b,
SystemHardLimitUSD: quota + token.UsedAmount,
})
}
func GetUsage(c *gin.Context) {
token := c.MustGet(ctxkey.Token).(*model.TokenCache)
c.JSON(http.StatusOK, openai.UsageResponse{TotalUsage: token.UsedAmount / 7 * 100})
token := middleware.GetToken(c)
c.JSON(http.StatusOK, openai.UsageResponse{TotalUsage: token.UsedAmount * 100})
}
+78 -18
View File
@@ -1,6 +1,7 @@
package controller
import (
"context"
"errors"
"fmt"
"io"
@@ -16,11 +17,12 @@ import (
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/common/helper"
"github.com/labring/sealos/service/aiproxy/common/render"
"github.com/labring/sealos/service/aiproxy/middleware"
"github.com/labring/sealos/service/aiproxy/model"
"github.com/labring/sealos/service/aiproxy/monitor"
"github.com/labring/sealos/service/aiproxy/relay/meta"
"github.com/labring/sealos/service/aiproxy/relay/relaymode"
"github.com/labring/sealos/service/aiproxy/relay/utils"
log "github.com/sirupsen/logrus"
)
@@ -28,8 +30,12 @@ import (
const channelTestRequestID = "channel-test"
// testSingleModel tests a single model in the channel
func testSingleModel(channel *model.Channel, modelName string) (*model.ChannelTest, error) {
body, mode, err := utils.BuildRequest(modelName)
func testSingleModel(mc *model.ModelCaches, channel *model.Channel, modelName string) (*model.ChannelTest, error) {
modelConfig, ok := mc.ModelConfig.GetModelConfig(modelName)
if !ok {
return nil, errors.New(modelName + " model config not found")
}
body, mode, err := utils.BuildRequest(modelConfig)
if err != nil {
return nil, err
}
@@ -37,38 +43,49 @@ func testSingleModel(channel *model.Channel, modelName string) (*model.ChannelTe
w := httptest.NewRecorder()
newc, _ := gin.CreateTestContext(w)
newc.Request = &http.Request{
Method: http.MethodPost,
URL: &url.URL{Path: utils.BuildModeDefaultPath(mode)},
URL: &url.URL{},
Body: io.NopCloser(body),
Header: make(http.Header),
}
newc.Set(string(helper.RequestIDKey), channelTestRequestID)
middleware.SetRequestID(newc, channelTestRequestID)
meta := meta.NewMeta(
channel,
mode,
modelName,
modelConfig,
meta.WithRequestID(channelTestRequestID),
meta.WithChannelTest(true),
)
bizErr := relayHelper(meta, newc)
relayController, ok := relayController(mode)
if !ok {
return nil, fmt.Errorf("relay mode %d not implemented", mode)
}
bizErr := relayController(meta, newc)
success := bizErr == nil
var respStr string
var code int
if bizErr == nil {
respStr = w.Body.String()
if success {
switch meta.Mode {
case relaymode.AudioSpeech,
relaymode.ImagesGenerations:
respStr = ""
default:
respStr = w.Body.String()
}
code = w.Code
} else {
respStr = bizErr.String()
respStr = bizErr.Error.JSONOrEmpty()
code = bizErr.StatusCode
}
return channel.UpdateModelTest(
meta.RequestAt,
meta.OriginModelName,
meta.ActualModelName,
meta.OriginModel,
meta.ActualModel,
meta.Mode,
time.Since(meta.RequestAt).Seconds(),
bizErr == nil,
success,
respStr,
code,
)
@@ -111,7 +128,7 @@ func TestChannel(c *gin.Context) {
return
}
ct, err := testSingleModel(channel, modelName)
ct, err := testSingleModel(model.LoadModelCaches(), channel, modelName)
if err != nil {
log.Errorf("failed to test channel %s(%d) model %s: %s", channel.Name, channel.ID, modelName, err.Error())
c.JSON(http.StatusOK, middleware.APIResponse{
@@ -137,8 +154,8 @@ type testResult struct {
Success bool `json:"success"`
}
func processTestResult(channel *model.Channel, modelName string, returnSuccess bool, successResponseBody bool) *testResult {
ct, err := testSingleModel(channel, modelName)
func processTestResult(mc *model.ModelCaches, channel *model.Channel, modelName string, returnSuccess bool, successResponseBody bool) *testResult {
ct, err := testSingleModel(mc, channel, modelName)
e := &utils.UnsupportedModelTypeError{}
if errors.As(err, &e) {
@@ -211,6 +228,8 @@ func TestChannelModels(c *gin.Context) {
models[i], models[j] = models[j], models[i]
})
mc := model.LoadModelCaches()
for _, modelName := range models {
wg.Add(1)
semaphore <- struct{}{}
@@ -219,7 +238,7 @@ func TestChannelModels(c *gin.Context) {
defer wg.Done()
defer func() { <-semaphore }()
result := processTestResult(channel, model, returnSuccess, successResponseBody)
result := processTestResult(mc, channel, model, returnSuccess, successResponseBody)
if result == nil {
return
}
@@ -294,6 +313,8 @@ func TestAllChannels(c *gin.Context) {
newChannels[i], newChannels[j] = newChannels[j], newChannels[i]
})
mc := model.LoadModelCaches()
for _, channel := range newChannels {
channelHasError := &atomic.Bool{}
hasErrorMap[channel.ID] = channelHasError
@@ -311,7 +332,7 @@ func TestAllChannels(c *gin.Context) {
defer wg.Done()
defer func() { <-semaphore }()
result := processTestResult(ch, model, returnSuccess, successResponseBody)
result := processTestResult(mc, ch, model, returnSuccess, successResponseBody)
if result == nil {
return
}
@@ -350,3 +371,42 @@ func TestAllChannels(c *gin.Context) {
})
}
}
func AutoTestBannedModels() {
log := log.WithFields(log.Fields{
"auto_test_banned_models": "true",
})
channels, err := monitor.GetAllBannedChannels(context.Background())
if err != nil {
log.Errorf("failed to get banned channels: %s", err.Error())
return
}
if len(channels) == 0 {
return
}
mc := model.LoadModelCaches()
for modelName, ids := range channels {
for _, id := range ids {
channel, err := model.LoadChannelByID(int(id))
if err != nil {
log.Errorf("failed to get channel by model %s: %s", modelName, err.Error())
continue
}
result, err := testSingleModel(mc, channel, modelName)
if err != nil {
log.Errorf("failed to test channel %s(%d) model %s: %s", channel.Name, channel.ID, modelName, err.Error())
}
if result.Success {
log.Infof("model %s(%d) test success, unban it", modelName, channel.ID)
err = monitor.ClearChannelModelErrors(context.Background(), modelName, channel.ID)
if err != nil {
log.Errorf("clear channel errors failed: %+v", err)
}
} else {
log.Infof("model %s(%d) test failed", modelName, channel.ID)
}
}
}
}
+70 -23
View File
@@ -1,6 +1,7 @@
package controller
import (
"fmt"
"maps"
"net/http"
"slices"
@@ -10,13 +11,20 @@ import (
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/middleware"
"github.com/labring/sealos/service/aiproxy/model"
"github.com/labring/sealos/service/aiproxy/monitor"
"github.com/labring/sealos/service/aiproxy/relay/adaptor"
"github.com/labring/sealos/service/aiproxy/relay/channeltype"
log "github.com/sirupsen/logrus"
)
func ChannelTypeNames(c *gin.Context) {
middleware.SuccessResponse(c, channeltype.ChannelNames)
}
func ChannelTypeMetas(c *gin.Context) {
middleware.SuccessResponse(c, channeltype.ChannelMetas)
}
func GetChannels(c *gin.Context) {
p, _ := strconv.Atoi(c.Query("p"))
p--
@@ -35,7 +43,7 @@ func GetChannels(c *gin.Context) {
channelType, _ := strconv.Atoi(c.Query("channel_type"))
baseURL := c.Query("base_url")
order := c.Query("order")
channels, total, err := model.GetChannels(p*perPage, perPage, false, false, id, name, key, channelType, baseURL, order)
channels, total, err := model.GetChannels(p*perPage, perPage, id, name, key, channelType, baseURL, order)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
@@ -47,7 +55,7 @@ func GetChannels(c *gin.Context) {
}
func GetAllChannels(c *gin.Context) {
channels, err := model.GetAllChannels(false, false)
channels, err := model.GetAllChannels()
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
@@ -64,7 +72,12 @@ func AddChannels(c *gin.Context) {
}
_channels := make([]*model.Channel, 0, len(channels))
for _, channel := range channels {
_channels = append(_channels, channel.ToChannels()...)
channels, err := channel.ToChannels()
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
_channels = append(_channels, channels...)
}
err = model.BatchInsertChannels(_channels)
if err != nil {
@@ -93,7 +106,7 @@ func SearchChannels(c *gin.Context) {
channelType, _ := strconv.Atoi(c.Query("channel_type"))
baseURL := c.Query("base_url")
order := c.Query("order")
channels, total, err := model.SearchChannels(keyword, p*perPage, perPage, false, false, id, name, key, channelType, baseURL, order)
channels, total, err := model.SearchChannels(keyword, p*perPage, perPage, id, name, key, channelType, baseURL, order)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
@@ -110,7 +123,7 @@ func GetChannel(c *gin.Context) {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
channel, err := model.GetChannelByID(id, false)
channel, err := model.GetChannelByID(id)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
@@ -119,19 +132,33 @@ func GetChannel(c *gin.Context) {
}
type AddChannelRequest struct {
ModelMapping map[string]string `json:"model_mapping"`
Config model.ChannelConfig `json:"config"`
Name string `json:"name"`
Key string `json:"key"`
BaseURL string `json:"base_url"`
Other string `json:"other"`
Models []string `json:"models"`
Type int `json:"type"`
Priority int32 `json:"priority"`
Status int `json:"status"`
ModelMapping map[string]string `json:"model_mapping"`
Config *model.ChannelConfig `json:"config"`
Name string `json:"name"`
Key string `json:"key"`
BaseURL string `json:"base_url"`
Other string `json:"other"`
Models []string `json:"models"`
Type int `json:"type"`
Priority int32 `json:"priority"`
Status int `json:"status"`
}
func (r *AddChannelRequest) ToChannel() *model.Channel {
func (r *AddChannelRequest) ToChannel() (*model.Channel, error) {
channelType, ok := channeltype.GetAdaptor(r.Type)
if !ok {
return nil, fmt.Errorf("invalid channel type: %d", r.Type)
}
if validator, ok := channelType.(adaptor.KeyValidator); ok {
err := validator.ValidateKey(r.Key)
if err != nil {
keyHelp := validator.KeyHelp()
if keyHelp == "" {
return nil, fmt.Errorf("%s [%s(%d)] invalid key: %w", r.Name, channeltype.ChannelNames[r.Type], r.Type, err)
}
return nil, fmt.Errorf("%s [%s(%d)] invalid key: %w, %s", r.Name, channeltype.ChannelNames[r.Type], r.Type, err, keyHelp)
}
}
return &model.Channel{
Type: r.Type,
Name: r.Name,
@@ -139,24 +166,27 @@ func (r *AddChannelRequest) ToChannel() *model.Channel {
BaseURL: r.BaseURL,
Models: slices.Clone(r.Models),
ModelMapping: maps.Clone(r.ModelMapping),
Config: r.Config,
Priority: r.Priority,
Status: r.Status,
}
Config: r.Config,
}, nil
}
func (r *AddChannelRequest) ToChannels() []*model.Channel {
func (r *AddChannelRequest) ToChannels() ([]*model.Channel, error) {
keys := strings.Split(r.Key, "\n")
channels := make([]*model.Channel, 0, len(keys))
for _, key := range keys {
if key == "" {
continue
}
c := r.ToChannel()
c, err := r.ToChannel()
if err != nil {
return nil, err
}
c.Key = key
channels = append(channels, c)
}
return channels
return channels, nil
}
func AddChannel(c *gin.Context) {
@@ -166,7 +196,12 @@ func AddChannel(c *gin.Context) {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
err = model.BatchInsertChannels(channel.ToChannels())
channels, err := channel.ToChannels()
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
err = model.BatchInsertChannels(channels)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
@@ -216,13 +251,21 @@ func UpdateChannel(c *gin.Context) {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
ch := channel.ToChannel()
ch, err := channel.ToChannel()
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
ch.ID = id
err = model.UpdateChannel(ch)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
err = monitor.ClearChannelAllModelErrors(c.Request.Context(), id)
if err != nil {
log.Errorf("failed to clear channel all model errors: %+v", err)
}
middleware.SuccessResponse(c, ch)
}
@@ -243,5 +286,9 @@ func UpdateChannelStatus(c *gin.Context) {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
err = monitor.ClearChannelAllModelErrors(c.Request.Context(), id)
if err != nil {
log.Errorf("failed to clear channel all model errors: %+v", err)
}
middleware.SuccessResponse(c, nil)
}
+217
View File
@@ -0,0 +1,217 @@
package controller
import (
"errors"
"fmt"
"net/http"
"strconv"
"time"
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/common/rpmlimit"
"github.com/labring/sealos/service/aiproxy/middleware"
"github.com/labring/sealos/service/aiproxy/model"
)
func getDashboardTime(t string) (time.Time, time.Time, time.Duration) {
end := time.Now()
var start time.Time
var timeSpan time.Duration
switch t {
case "month":
start = end.AddDate(0, 0, -30)
timeSpan = time.Hour * 24
case "two_week":
start = end.AddDate(0, 0, -15)
timeSpan = time.Hour * 24
case "week":
start = end.AddDate(0, 0, -7)
timeSpan = time.Hour * 24
case "day":
fallthrough
default:
start = end.AddDate(0, 0, -1)
timeSpan = time.Hour * 1
}
return start, end, timeSpan
}
func fillGaps(data []*model.ChartData, start, end time.Time, timeSpan time.Duration) []*model.ChartData {
if len(data) == 0 {
return data
}
// Handle first point
firstPoint := time.Unix(data[0].Timestamp, 0)
firstAlignedTime := firstPoint
for !firstAlignedTime.Add(-timeSpan).Before(start) {
firstAlignedTime = firstAlignedTime.Add(-timeSpan)
}
var firstIsZero bool
if !firstAlignedTime.Equal(firstPoint) {
data = append([]*model.ChartData{
{
Timestamp: firstAlignedTime.Unix(),
},
}, data...)
firstIsZero = true
}
// Handle last point
lastPoint := time.Unix(data[len(data)-1].Timestamp, 0)
lastAlignedTime := lastPoint
for !lastAlignedTime.Add(timeSpan).After(end) {
lastAlignedTime = lastAlignedTime.Add(timeSpan)
}
var lastIsZero bool
if !lastAlignedTime.Equal(lastPoint) {
data = append(data, &model.ChartData{
Timestamp: lastAlignedTime.Unix(),
})
lastIsZero = true
}
result := make([]*model.ChartData, 0, len(data))
result = append(result, data[0])
for i := 1; i < len(data); i++ {
curr := data[i]
prev := data[i-1]
hourDiff := (curr.Timestamp - prev.Timestamp) / int64(timeSpan.Seconds())
// If gap is 1 hour or less, continue
if hourDiff <= 1 {
result = append(result, curr)
continue
}
// If gap is more than 3 hours, only add boundary points
if hourDiff > 3 {
// Add point for hour after prev
if i != 1 || (i == 1 && !firstIsZero) {
result = append(result, &model.ChartData{
Timestamp: prev.Timestamp + int64(timeSpan.Seconds()),
})
}
// Add point for hour before curr
if i != len(data)-1 || (i == len(data)-1 && !lastIsZero) {
result = append(result, &model.ChartData{
Timestamp: curr.Timestamp - int64(timeSpan.Seconds()),
})
}
result = append(result, curr)
continue
}
// Fill gaps of 2-3 hours with zero points
for j := prev.Timestamp + int64(timeSpan.Seconds()); j < curr.Timestamp; j += int64(timeSpan.Seconds()) {
result = append(result, &model.ChartData{
Timestamp: j,
})
}
result = append(result, curr)
}
return result
}
func getTimeSpanWithDefault(c *gin.Context, defaultTimeSpan time.Duration) time.Duration {
spanStr := c.Query("span")
if spanStr == "" {
return defaultTimeSpan
}
span, err := strconv.Atoi(spanStr)
if err != nil {
return defaultTimeSpan
}
if span < 1 || span > 48 {
return defaultTimeSpan
}
return time.Duration(span) * time.Hour
}
func GetDashboard(c *gin.Context) {
log := middleware.GetLogger(c)
start, end, timeSpan := getDashboardTime(c.Query("type"))
modelName := c.Query("model")
timeSpan = getTimeSpanWithDefault(c, timeSpan)
dashboards, err := model.GetDashboardData(start, end, modelName, timeSpan)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
dashboards.ChartData = fillGaps(dashboards.ChartData, start, end, timeSpan)
if common.RedisEnabled {
rpm, err := rpmlimit.GetRPM(c.Request.Context(), "", modelName)
if err != nil {
log.Errorf("failed to get rpm: %v", err)
} else {
dashboards.RPM = rpm
}
}
middleware.SuccessResponse(c, dashboards)
}
func GetGroupDashboard(c *gin.Context) {
log := middleware.GetLogger(c)
group := c.Param("group")
if group == "" {
middleware.ErrorResponse(c, http.StatusOK, "invalid parameter")
return
}
start, end, timeSpan := getDashboardTime(c.Query("type"))
tokenName := c.Query("token_name")
modelName := c.Query("model")
timeSpan = getTimeSpanWithDefault(c, timeSpan)
dashboards, err := model.GetGroupDashboardData(group, start, end, tokenName, modelName, timeSpan)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, "failed to get statistics")
return
}
dashboards.ChartData = fillGaps(dashboards.ChartData, start, end, timeSpan)
if common.RedisEnabled && tokenName == "" {
rpm, err := rpmlimit.GetRPM(c.Request.Context(), group, modelName)
if err != nil {
log.Errorf("failed to get rpm: %v", err)
} else {
dashboards.RPM = rpm
}
}
middleware.SuccessResponse(c, dashboards)
}
func GetGroupDashboardModels(c *gin.Context) {
group := c.Param("group")
if group == "" {
middleware.ErrorResponse(c, http.StatusOK, "invalid parameter")
return
}
groupCache, err := model.CacheGetGroup(group)
if err != nil {
if errors.Is(err, model.NotFoundError(model.ErrGroupNotFound)) {
middleware.SuccessResponse(c, model.LoadModelCaches().EnabledModelConfigs)
} else {
middleware.ErrorResponse(c, http.StatusOK, fmt.Sprintf("failed to get group: %v", err))
}
return
}
enabledModelConfigs := model.LoadModelCaches().EnabledModelConfigs
newEnabledModelConfigs := make([]*model.ModelConfig, len(enabledModelConfigs))
for i, mc := range enabledModelConfigs {
newEnabledModelConfigs[i] = middleware.GetGroupAdjustedModelConfig(groupCache, mc)
}
middleware.SuccessResponse(c, newEnabledModelConfigs)
}
+147 -42
View File
@@ -5,14 +5,30 @@ import (
"strconv"
"time"
"github.com/gin-gonic/gin"
json "github.com/json-iterator/go"
"github.com/labring/sealos/service/aiproxy/middleware"
"github.com/labring/sealos/service/aiproxy/model"
"github.com/gin-gonic/gin"
)
type GroupResponse struct {
*model.Group
AccessedAt time.Time `json:"accessed_at,omitempty"`
}
func (g *GroupResponse) MarshalJSON() ([]byte, error) {
type Alias model.Group
return json.Marshal(&struct {
*Alias
CreatedAt int64 `json:"created_at,omitempty"`
AccessedAt int64 `json:"accessed_at,omitempty"`
}{
Alias: (*Alias)(g.Group),
CreatedAt: g.CreatedAt.UnixMilli(),
AccessedAt: g.AccessedAt.UnixMilli(),
})
}
func GetGroups(c *gin.Context) {
p, _ := strconv.Atoi(c.Query("p"))
p--
@@ -32,8 +48,16 @@ func GetGroups(c *gin.Context) {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
groupResponses := make([]*GroupResponse, len(groups))
for i, group := range groups {
lastRequestAt, _ := model.GetGroupLastRequestTime(group.ID)
groupResponses[i] = &GroupResponse{
Group: group,
AccessedAt: lastRequestAt,
}
}
middleware.SuccessResponse(c, gin.H{
"groups": groups,
"groups": groupResponses,
"total": total,
})
}
@@ -58,57 +82,128 @@ func SearchGroups(c *gin.Context) {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
groupResponses := make([]*GroupResponse, len(groups))
for i, group := range groups {
lastRequestAt, _ := model.GetGroupLastRequestTime(group.ID)
groupResponses[i] = &GroupResponse{
Group: group,
AccessedAt: lastRequestAt,
}
}
middleware.SuccessResponse(c, gin.H{
"groups": groups,
"groups": groupResponses,
"total": total,
})
}
func GetGroup(c *gin.Context) {
id := c.Param("id")
if id == "" {
group := c.Param("group")
if group == "" {
middleware.ErrorResponse(c, http.StatusOK, "group id is empty")
return
}
group, err := model.GetGroupByID(id)
_group, err := model.GetGroupByID(group)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
middleware.SuccessResponse(c, group)
}
func GetGroupDashboard(c *gin.Context) {
id := c.Param("id")
now := time.Now()
startOfDay := now.Truncate(24*time.Hour).AddDate(0, 0, -6).Unix()
endOfDay := now.Truncate(24 * time.Hour).Add(24*time.Hour - time.Second).Unix()
dashboards, err := model.SearchLogsByDayAndModel(id, time.Unix(startOfDay, 0), time.Unix(endOfDay, 0))
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, "failed to get statistics")
return
lastRequestAt, _ := model.GetGroupLastRequestTime(group)
groupResponse := &GroupResponse{
Group: _group,
AccessedAt: lastRequestAt,
}
middleware.SuccessResponse(c, dashboards)
middleware.SuccessResponse(c, groupResponse)
}
type UpdateGroupQPMRequest struct {
QPM int64 `json:"qpm"`
type UpdateGroupRPMRatioRequest struct {
RPMRatio float64 `json:"rpm_ratio"`
}
func UpdateGroupQPM(c *gin.Context) {
id := c.Param("id")
if id == "" {
func UpdateGroupRPMRatio(c *gin.Context) {
group := c.Param("group")
if group == "" {
middleware.ErrorResponse(c, http.StatusOK, "invalid parameter")
return
}
req := UpdateGroupQPMRequest{}
req := UpdateGroupRPMRatioRequest{}
err := json.NewDecoder(c.Request.Body).Decode(&req)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, "invalid parameter")
return
}
err = model.UpdateGroupQPM(id, req.QPM)
err = model.UpdateGroupRPMRatio(group, req.RPMRatio)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
middleware.SuccessResponse(c, nil)
}
type UpdateGroupRPMRequest struct {
RPM map[string]int64 `json:"rpm"`
}
func UpdateGroupRPM(c *gin.Context) {
group := c.Param("group")
if group == "" {
middleware.ErrorResponse(c, http.StatusOK, "invalid parameter")
return
}
req := UpdateGroupRPMRequest{}
err := json.NewDecoder(c.Request.Body).Decode(&req)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, "invalid parameter")
return
}
err = model.UpdateGroupRPM(group, req.RPM)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
middleware.SuccessResponse(c, nil)
}
type UpdateGroupTPMRequest struct {
TPM map[string]int64 `json:"tpm"`
}
func UpdateGroupTPM(c *gin.Context) {
group := c.Param("group")
if group == "" {
middleware.ErrorResponse(c, http.StatusOK, "invalid parameter")
return
}
req := UpdateGroupTPMRequest{}
err := json.NewDecoder(c.Request.Body).Decode(&req)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, "invalid parameter")
return
}
err = model.UpdateGroupTPM(group, req.TPM)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
middleware.SuccessResponse(c, nil)
}
type UpdateGroupTPMRatioRequest struct {
TPMRatio float64 `json:"tpm_ratio"`
}
func UpdateGroupTPMRatio(c *gin.Context) {
group := c.Param("group")
if group == "" {
middleware.ErrorResponse(c, http.StatusOK, "invalid parameter")
return
}
req := UpdateGroupTPMRatioRequest{}
err := json.NewDecoder(c.Request.Body).Decode(&req)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, "invalid parameter")
return
}
err = model.UpdateGroupTPMRatio(group, req.TPMRatio)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
@@ -121,8 +216,8 @@ type UpdateGroupStatusRequest struct {
}
func UpdateGroupStatus(c *gin.Context) {
id := c.Param("id")
if id == "" {
group := c.Param("group")
if group == "" {
middleware.ErrorResponse(c, http.StatusOK, "invalid parameter")
return
}
@@ -132,7 +227,7 @@ func UpdateGroupStatus(c *gin.Context) {
middleware.ErrorResponse(c, http.StatusOK, "invalid parameter")
return
}
err = model.UpdateGroupStatus(id, req.Status)
err = model.UpdateGroupStatus(group, req.Status)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
@@ -141,12 +236,12 @@ func UpdateGroupStatus(c *gin.Context) {
}
func DeleteGroup(c *gin.Context) {
id := c.Param("id")
if id == "" {
group := c.Param("group")
if group == "" {
middleware.ErrorResponse(c, http.StatusOK, "invalid parameter")
return
}
err := model.DeleteGroupByID(id)
err := model.DeleteGroupByID(group)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
@@ -170,20 +265,30 @@ func DeleteGroups(c *gin.Context) {
}
type CreateGroupRequest struct {
ID string `json:"id"`
QPM int64 `json:"qpm"`
RPM map[string]int64 `json:"rpm"`
RPMRatio float64 `json:"rpm_ratio"`
TPM map[string]int64 `json:"tpm"`
TPMRatio float64 `json:"tpm_ratio"`
}
func CreateGroup(c *gin.Context) {
var group CreateGroupRequest
err := json.NewDecoder(c.Request.Body).Decode(&group)
if err != nil || group.ID == "" {
group := c.Param("group")
if group == "" {
middleware.ErrorResponse(c, http.StatusOK, "invalid parameter")
return
}
req := CreateGroupRequest{}
err := json.NewDecoder(c.Request.Body).Decode(&req)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, "invalid parameter")
return
}
if err := model.CreateGroup(&model.Group{
ID: group.ID,
QPM: group.QPM,
ID: group,
RPMRatio: req.RPMRatio,
RPM: req.RPM,
TPMRatio: req.TPMRatio,
TPM: req.TPM,
}); err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
+127 -98
View File
@@ -22,7 +22,6 @@ func GetLogs(c *gin.Context) {
} else if perPage > 100 {
perPage = 100
}
code, _ := strconv.Atoi(c.Query("code"))
startTimestamp, _ := strconv.ParseInt(c.Query("start_timestamp"), 10, 64)
endTimestamp, _ := strconv.ParseInt(c.Query("end_timestamp"), 10, 64)
var startTimestampTime time.Time
@@ -35,42 +34,47 @@ func GetLogs(c *gin.Context) {
}
tokenName := c.Query("token_name")
modelName := c.Query("model_name")
channel, _ := strconv.Atoi(c.Query("channel"))
channelID, _ := strconv.Atoi(c.Query("channel"))
group := c.Query("group")
endpoint := c.Query("endpoint")
content := c.Query("content")
tokenID, _ := strconv.Atoi(c.Query("token_id"))
order := c.Query("order")
requestID := c.Query("request_id")
mode, _ := strconv.Atoi(c.Query("mode"))
logs, total, err := model.GetLogs(
codeType := c.Query("code_type")
withBody, _ := strconv.ParseBool(c.Query("with_body"))
ip := c.Query("ip")
result, err := model.GetLogs(
group,
startTimestampTime,
endTimestampTime,
code,
modelName,
group,
requestID,
tokenID,
tokenName,
p*perPage,
perPage,
channel,
channelID,
endpoint,
content,
order,
mode,
model.CodeType(codeType),
withBody,
ip,
p,
perPage,
)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
middleware.SuccessResponse(c, gin.H{
"logs": logs,
"total": total,
})
middleware.SuccessResponse(c, result)
}
func GetGroupLogs(c *gin.Context) {
group := c.Param("group")
if group == "" {
middleware.ErrorResponse(c, http.StatusOK, "group is required")
return
}
p, _ := strconv.Atoi(c.Query("p"))
p--
if p < 0 {
@@ -82,7 +86,6 @@ func GetGroupLogs(c *gin.Context) {
} else if perPage > 100 {
perPage = 100
}
code, _ := strconv.Atoi(c.Query("code"))
startTimestamp, _ := strconv.ParseInt(c.Query("start_timestamp"), 10, 64)
endTimestamp, _ := strconv.ParseInt(c.Query("end_timestamp"), 10, 64)
var startTimestampTime time.Time
@@ -95,39 +98,38 @@ func GetGroupLogs(c *gin.Context) {
}
tokenName := c.Query("token_name")
modelName := c.Query("model_name")
channel, _ := strconv.Atoi(c.Query("channel"))
group := c.Param("group")
channelID, _ := strconv.Atoi(c.Query("channel"))
endpoint := c.Query("endpoint")
content := c.Query("content")
tokenID, _ := strconv.Atoi(c.Query("token_id"))
order := c.Query("order")
requestID := c.Query("request_id")
mode, _ := strconv.Atoi(c.Query("mode"))
logs, total, err := model.GetGroupLogs(
codeType := c.Query("code_type")
withBody, _ := strconv.ParseBool(c.Query("with_body"))
ip := c.Query("ip")
result, err := model.GetGroupLogs(
group,
startTimestampTime,
endTimestampTime,
code,
modelName,
requestID,
tokenID,
tokenName,
p*perPage,
perPage,
channel,
channelID,
endpoint,
content,
order,
mode,
model.CodeType(codeType),
withBody,
ip,
p,
perPage,
)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
middleware.SuccessResponse(c, gin.H{
"logs": logs,
"total": total,
})
middleware.SuccessResponse(c, result)
}
func SearchLogs(c *gin.Context) {
@@ -139,70 +141,10 @@ func SearchLogs(c *gin.Context) {
} else if perPage > 100 {
perPage = 100
}
code, _ := strconv.Atoi(c.Query("code"))
endpoint := c.Query("endpoint")
tokenName := c.Query("token_name")
modelName := c.Query("model_name")
content := c.Query("content")
groupID := c.Query("group_id")
tokenID, _ := strconv.Atoi(c.Query("token_id"))
channel, _ := strconv.Atoi(c.Query("channel"))
startTimestamp, _ := strconv.ParseInt(c.Query("start_timestamp"), 10, 64)
endTimestamp, _ := strconv.ParseInt(c.Query("end_timestamp"), 10, 64)
var startTimestampTime time.Time
if startTimestamp != 0 {
startTimestampTime = time.UnixMilli(startTimestamp)
}
var endTimestampTime time.Time
if endTimestamp != 0 {
endTimestampTime = time.UnixMilli(endTimestamp)
}
order := c.Query("order")
requestID := c.Query("request_id")
mode, _ := strconv.Atoi(c.Query("mode"))
logs, total, err := model.SearchLogs(
keyword,
p,
perPage,
code,
endpoint,
groupID,
requestID,
tokenID,
tokenName,
modelName,
content,
startTimestampTime,
endTimestampTime,
channel,
order,
mode,
)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
middleware.SuccessResponse(c, gin.H{
"logs": logs,
"total": total,
})
}
func SearchGroupLogs(c *gin.Context) {
keyword := c.Query("keyword")
p, _ := strconv.Atoi(c.Query("p"))
perPage, _ := strconv.Atoi(c.Query("per_page"))
if perPage <= 0 {
perPage = 10
} else if perPage > 100 {
perPage = 100
}
group := c.Param("group")
code, _ := strconv.Atoi(c.Query("code"))
endpoint := c.Query("endpoint")
tokenName := c.Query("token_name")
modelName := c.Query("model_name")
content := c.Query("content")
group := c.Query("group_id")
tokenID, _ := strconv.Atoi(c.Query("token_id"))
channelID, _ := strconv.Atoi(c.Query("channel"))
startTimestamp, _ := strconv.ParseInt(c.Query("start_timestamp"), 10, 64)
@@ -218,32 +160,119 @@ func SearchGroupLogs(c *gin.Context) {
order := c.Query("order")
requestID := c.Query("request_id")
mode, _ := strconv.Atoi(c.Query("mode"))
logs, total, err := model.SearchGroupLogs(
codeType := c.Query("code_type")
withBody, _ := strconv.ParseBool(c.Query("with_body"))
ip := c.Query("ip")
result, err := model.SearchLogs(
group,
keyword,
p,
perPage,
code,
endpoint,
requestID,
tokenID,
tokenName,
modelName,
content,
startTimestampTime,
endTimestampTime,
channelID,
order,
mode,
model.CodeType(codeType),
withBody,
ip,
p,
perPage,
)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
middleware.SuccessResponse(c, gin.H{
"logs": logs,
"total": total,
})
middleware.SuccessResponse(c, result)
}
func SearchGroupLogs(c *gin.Context) {
group := c.Param("group")
if group == "" {
middleware.ErrorResponse(c, http.StatusOK, "group is required")
return
}
keyword := c.Query("keyword")
p, _ := strconv.Atoi(c.Query("p"))
perPage, _ := strconv.Atoi(c.Query("per_page"))
if perPage <= 0 {
perPage = 10
} else if perPage > 100 {
perPage = 100
}
endpoint := c.Query("endpoint")
tokenName := c.Query("token_name")
modelName := c.Query("model_name")
tokenID, _ := strconv.Atoi(c.Query("token_id"))
channelID, _ := strconv.Atoi(c.Query("channel"))
startTimestamp, _ := strconv.ParseInt(c.Query("start_timestamp"), 10, 64)
endTimestamp, _ := strconv.ParseInt(c.Query("end_timestamp"), 10, 64)
var startTimestampTime time.Time
if startTimestamp != 0 {
startTimestampTime = time.UnixMilli(startTimestamp)
}
var endTimestampTime time.Time
if endTimestamp != 0 {
endTimestampTime = time.UnixMilli(endTimestamp)
}
order := c.Query("order")
requestID := c.Query("request_id")
mode, _ := strconv.Atoi(c.Query("mode"))
codeType := c.Query("code_type")
withBody, _ := strconv.ParseBool(c.Query("with_body"))
ip := c.Query("ip")
result, err := model.SearchGroupLogs(
group,
keyword,
endpoint,
requestID,
tokenID,
tokenName,
modelName,
startTimestampTime,
endTimestampTime,
channelID,
order,
mode,
model.CodeType(codeType),
withBody,
ip,
p,
perPage,
)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
middleware.SuccessResponse(c, result)
}
func GetLogDetail(c *gin.Context) {
logID, _ := strconv.Atoi(c.Param("log_id"))
log, err := model.GetLogDetail(logID)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
middleware.SuccessResponse(c, log)
}
func GetGroupLogDetail(c *gin.Context) {
group := c.Param("group")
if group == "" {
middleware.ErrorResponse(c, http.StatusOK, "group is required")
return
}
logID, _ := strconv.Atoi(c.Param("log_id"))
log, err := model.GetGroupLogDetail(logID, group)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
middleware.SuccessResponse(c, log)
}
func DeleteHistoryLogs(c *gin.Context) {
+1 -2
View File
@@ -1,10 +1,9 @@
package controller
import (
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/middleware"
"github.com/gin-gonic/gin"
)
type StatusData struct {
+28 -20
View File
@@ -168,11 +168,11 @@ func ChannelDefaultModelsAndMappingByType(c *gin.Context) {
}
func EnabledModels(c *gin.Context) {
middleware.SuccessResponse(c, model.CacheGetEnabledModelConfigs())
middleware.SuccessResponse(c, model.LoadModelCaches().EnabledModelConfigs)
}
func ChannelEnabledModels(c *gin.Context) {
middleware.SuccessResponse(c, model.CacheGetEnabledChannelType2ModelConfigs())
middleware.SuccessResponse(c, model.LoadModelCaches().EnabledChannelType2ModelConfigs)
}
func ChannelEnabledModelsByType(c *gin.Context) {
@@ -186,23 +186,26 @@ func ChannelEnabledModelsByType(c *gin.Context) {
middleware.ErrorResponse(c, http.StatusOK, "invalid type")
return
}
middleware.SuccessResponse(c, model.CacheGetEnabledChannelType2ModelConfigs()[channelTypeInt])
middleware.SuccessResponse(c, model.LoadModelCaches().EnabledChannelType2ModelConfigs[channelTypeInt])
}
func ListModels(c *gin.Context) {
models := model.CacheGetEnabledModelConfigs()
enabledModelConfigsMap := middleware.GetModelCaches(c).EnabledModelConfigsMap
token := middleware.GetToken(c)
availableOpenAIModels := make([]*OpenAIModels, len(models))
availableOpenAIModels := make([]*OpenAIModels, 0, len(token.Models))
for idx, model := range models {
availableOpenAIModels[idx] = &OpenAIModels{
ID: model.Model,
Object: "model",
Created: 1626777600,
OwnedBy: string(model.Owner),
Root: model.Model,
Permission: permission,
Parent: nil,
for _, model := range token.Models {
if mc, ok := enabledModelConfigsMap[model]; ok {
availableOpenAIModels = append(availableOpenAIModels, &OpenAIModels{
ID: model,
Object: "model",
Created: 1626777600,
OwnedBy: string(mc.Owner),
Root: model,
Permission: permission,
Parent: nil,
})
}
}
@@ -214,10 +217,15 @@ func ListModels(c *gin.Context) {
func RetrieveModel(c *gin.Context) {
modelName := c.Param("model")
enabledModels := model.GetEnabledModel2Channels()
model, ok := model.CacheGetModelConfig(modelName)
enabledModelConfigsMap := middleware.GetModelCaches(c).EnabledModelConfigsMap
if _, exist := enabledModels[modelName]; !exist || !ok {
mc, ok := enabledModelConfigsMap[modelName]
if ok {
token := middleware.GetToken(c)
ok = slices.Contains(token.Models, modelName)
}
if !ok {
c.JSON(200, gin.H{
"error": &relaymodel.Error{
Message: fmt.Sprintf("the model '%s' does not exist", modelName),
@@ -230,11 +238,11 @@ func RetrieveModel(c *gin.Context) {
}
c.JSON(200, &OpenAIModels{
ID: model.Model,
ID: modelName,
Object: "model",
Created: 1626777600,
OwnedBy: string(model.Owner),
Root: model.Model,
OwnedBy: string(mc.Owner),
Root: modelName,
Permission: permission,
Parent: nil,
})
+14 -4
View File
@@ -87,13 +87,23 @@ func SearchModelConfigs(c *gin.Context) {
})
}
type SaveModelConfigsRequest struct {
CreatedAt int64 `json:"created_at"`
UpdatedAt int64 `json:"updated_at"`
*model.ModelConfig
}
func SaveModelConfigs(c *gin.Context) {
var configs []*model.ModelConfig
var configs []*SaveModelConfigsRequest
if err := c.ShouldBindJSON(&configs); err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
err := model.SaveModelConfigs(configs)
modelConfigs := make([]*model.ModelConfig, len(configs))
for i, config := range configs {
modelConfigs[i] = config.ModelConfig
}
err := model.SaveModelConfigs(modelConfigs)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
@@ -102,12 +112,12 @@ func SaveModelConfigs(c *gin.Context) {
}
func SaveModelConfig(c *gin.Context) {
var config model.ModelConfig
var config SaveModelConfigsRequest
if err := c.ShouldBindJSON(&config); err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
err := model.SaveModelConfig(&config)
err := model.SaveModelConfig(config.ModelConfig)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
+74
View File
@@ -0,0 +1,74 @@
package controller
import (
"net/http"
"strconv"
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/middleware"
"github.com/labring/sealos/service/aiproxy/monitor"
)
func GetAllChannelModelErrorRates(c *gin.Context) {
rates, err := monitor.GetAllChannelModelErrorRates(c.Request.Context())
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
c.JSON(http.StatusOK, rates)
}
func GetChannelModelErrorRates(c *gin.Context) {
channelID := c.Param("id")
channelIDInt, err := strconv.ParseInt(channelID, 10, 64)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, "Invalid channel ID")
return
}
rates, err := monitor.GetChannelModelErrorRates(c.Request.Context(), channelIDInt)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
c.JSON(http.StatusOK, rates)
}
func ClearAllModelErrors(c *gin.Context) {
err := monitor.ClearAllModelErrors(c.Request.Context())
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
c.Status(http.StatusNoContent)
}
func ClearChannelAllModelErrors(c *gin.Context) {
channelID := c.Param("id")
channelIDInt, err := strconv.ParseInt(channelID, 10, 64)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, "Invalid channel ID")
return
}
err = monitor.ClearChannelAllModelErrors(c.Request.Context(), int(channelIDInt))
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
c.Status(http.StatusNoContent)
}
func ClearChannelModelErrors(c *gin.Context) {
channelID := c.Param("id")
model := c.Param("model")
channelIDInt, err := strconv.ParseInt(channelID, 10, 64)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, "Invalid channel ID")
return
}
err = monitor.ClearChannelModelErrors(c.Request.Context(), model, int(channelIDInt))
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
c.Status(http.StatusNoContent)
}
+15 -3
View File
@@ -3,12 +3,10 @@ package controller
import (
"net/http"
"github.com/gin-gonic/gin"
json "github.com/json-iterator/go"
"github.com/labring/sealos/service/aiproxy/middleware"
"github.com/labring/sealos/service/aiproxy/model"
"github.com/gin-gonic/gin"
)
func GetOptions(c *gin.Context) {
@@ -24,6 +22,20 @@ func GetOptions(c *gin.Context) {
middleware.SuccessResponse(c, options)
}
func GetOption(c *gin.Context) {
key := c.Param("key")
if key == "" {
middleware.ErrorResponse(c, http.StatusOK, "key is required")
return
}
option, err := model.GetOption(key)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
middleware.SuccessResponse(c, option)
}
func UpdateOption(c *gin.Context) {
var option model.Option
err := json.NewDecoder(c.Request.Body).Decode(&option)
+134 -58
View File
@@ -2,112 +2,188 @@ package controller
import (
"bytes"
"context"
"errors"
"io"
"net/http"
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/common/config"
"github.com/labring/sealos/service/aiproxy/common/helper"
"github.com/labring/sealos/service/aiproxy/middleware"
dbmodel "github.com/labring/sealos/service/aiproxy/model"
"github.com/labring/sealos/service/aiproxy/monitor"
"github.com/labring/sealos/service/aiproxy/relay/controller"
"github.com/labring/sealos/service/aiproxy/relay/meta"
"github.com/labring/sealos/service/aiproxy/relay/model"
"github.com/labring/sealos/service/aiproxy/relay/relaymode"
log "github.com/sirupsen/logrus"
)
// https://platform.openai.com/docs/api-reference/chat
func relayHelper(meta *meta.Meta, c *gin.Context) *model.ErrorWithStatusCode {
log := middleware.GetLogger(c)
middleware.SetLogFieldsFromMeta(meta, log.Data)
switch meta.Mode {
case relaymode.ImagesGenerations:
return controller.RelayImageHelper(meta, c)
type RelayController func(*meta.Meta, *gin.Context) *model.ErrorWithStatusCode
func relayController(mode int) (RelayController, bool) {
var relayController RelayController
switch mode {
case relaymode.ImagesGenerations,
relaymode.Edits:
relayController = controller.RelayImageHelper
case relaymode.AudioSpeech:
return controller.RelayTTSHelper(meta, c)
case relaymode.AudioTranslation:
return controller.RelaySTTHelper(meta, c)
case relaymode.AudioTranscription:
return controller.RelaySTTHelper(meta, c)
relayController = controller.RelayTTSHelper
case relaymode.AudioTranslation,
relaymode.AudioTranscription:
relayController = controller.RelaySTTHelper
case relaymode.Rerank:
return controller.RerankHelper(meta, c)
relayController = controller.RerankHelper
case relaymode.ChatCompletions,
relaymode.Embeddings,
relaymode.Completions,
relaymode.Moderations:
relayController = controller.RelayTextHelper
default:
return controller.RelayTextHelper(meta, c)
return nil, false
}
return func(meta *meta.Meta, c *gin.Context) *model.ErrorWithStatusCode {
log := middleware.GetLogger(c)
middleware.SetLogFieldsFromMeta(meta, log.Data)
return relayController(meta, c)
}, true
}
func RelayHelper(meta *meta.Meta, c *gin.Context, relayController RelayController) (*model.ErrorWithStatusCode, bool) {
err := relayController(meta, c)
if err == nil {
if err := monitor.AddRequest(
context.Background(),
meta.OriginModel,
int64(meta.Channel.ID),
false,
); err != nil {
log.Errorf("add request failed: %+v", err)
}
return nil, false
}
if shouldRetry(c, err.StatusCode) {
if err := monitor.AddRequest(
context.Background(),
meta.OriginModel,
int64(meta.Channel.ID),
true,
); err != nil {
log.Errorf("add request failed: %+v", err)
}
return err, true
}
return err, false
}
func getChannelWithFallback(cache *dbmodel.ModelCaches, model string, failedChannelIDs ...int) (*dbmodel.Channel, error) {
channel, err := cache.GetRandomSatisfiedChannel(model, failedChannelIDs...)
if err == nil {
return channel, nil
}
if !errors.Is(err, dbmodel.ErrChannelsExhausted) {
return nil, err
}
return cache.GetRandomSatisfiedChannel(model)
}
func NewRelay(mode int) func(c *gin.Context) {
relayController, ok := relayController(mode)
if !ok {
log.Fatalf("relay mode %d not implemented", mode)
}
return func(c *gin.Context) {
relay(c, mode, relayController)
}
}
func Relay(c *gin.Context) {
func relay(c *gin.Context, mode int, relayController RelayController) {
log := middleware.GetLogger(c)
if config.DebugEnabled {
requestBody, _ := common.GetRequestBody(c.Request)
log.Debugf("request body: %s", requestBody)
requestModel := middleware.GetOriginalModel(c)
ids, err := monitor.GetBannedChannels(c.Request.Context(), requestModel)
if err != nil {
log.Errorf("get %s auto banned channels failed: %+v", requestModel, err)
}
meta := middleware.NewMetaByContext(c)
bizErr := relayHelper(meta, c)
log.Debugf("%s model banned channels: %+v", requestModel, ids)
failedChannelIDs := []int{}
for _, id := range ids {
failedChannelIDs = append(failedChannelIDs, int(id))
}
mc := middleware.GetModelCaches(c)
channel, err := getChannelWithFallback(mc, requestModel, failedChannelIDs...)
if err != nil {
c.JSON(http.StatusServiceUnavailable, gin.H{
"error": &model.Error{
Message: "The upstream load is saturated, please try again later",
Code: "upstream_load_saturated",
Type: middleware.ErrorTypeAIPROXY,
},
})
return
}
meta := middleware.NewMetaByContext(c, channel, requestModel, mode)
bizErr, retry := RelayHelper(meta, c, relayController)
if bizErr == nil {
return
}
lastFailedChannelID := meta.Channel.ID
requestID := c.GetString(string(helper.RequestIDKey))
retryTimes := config.GetRetryTimes()
if !shouldRetry(c, bizErr.StatusCode) {
retryTimes = 0
failedChannelIDs = append(failedChannelIDs, channel.ID)
requestID := middleware.GetRequestID(c)
var retryTimes int64
if retry {
retryTimes = config.GetRetryTimes()
}
for i := retryTimes; i > 0; i-- {
channel, err := dbmodel.CacheGetRandomSatisfiedChannel(meta.OriginModelName)
newChannel, err := mc.GetRandomSatisfiedChannel(requestModel, failedChannelIDs...)
if err != nil {
log.Errorf("get random satisfied channel failed: %+v", err)
break
}
log.Infof("using channel #%d to retry (remain times %d)", channel.ID, i)
if channel.ID == lastFailedChannelID {
continue
if errors.Is(err, dbmodel.ErrChannelsNotFound) {
break
}
if !errors.Is(err, dbmodel.ErrChannelsExhausted) {
break
}
newChannel = channel
}
log.Warnf("using channel %s(%d) to retry (remain times %d)", newChannel.Name, newChannel.ID, i)
requestBody, err := common.GetRequestBody(c.Request)
if err != nil {
log.Errorf("GetRequestBody failed: %+v", err)
break
}
c.Request.Body = io.NopCloser(bytes.NewBuffer(requestBody))
meta.Reset(channel)
bizErr = relayHelper(meta, c)
meta.Reset(newChannel)
bizErr, retry = RelayHelper(meta, c, relayController)
if bizErr == nil {
return
}
lastFailedChannelID = channel.ID
if !retry {
break
}
failedChannelIDs = append(failedChannelIDs, newChannel.ID)
}
if bizErr != nil {
message := bizErr.Message
if bizErr.StatusCode == http.StatusTooManyRequests {
message = "The upstream load of the current group is saturated, please try again later"
}
c.JSON(bizErr.StatusCode, gin.H{
"error": &model.Error{
Message: helper.MessageWithRequestID(message, requestID),
Code: bizErr.Code,
Param: bizErr.Param,
Type: bizErr.Type,
},
})
bizErr.Error.Message = middleware.MessageWithRequestID(bizErr.Error.Message, requestID)
c.JSON(bizErr.StatusCode, bizErr)
}
}
// 仅当是channel错误时,才需要重试,用户请求参数错误时,不需要重试
func shouldRetry(_ *gin.Context, statusCode int) bool {
if statusCode == http.StatusTooManyRequests {
if statusCode == http.StatusTooManyRequests ||
statusCode == http.StatusGatewayTimeout ||
statusCode == http.StatusForbidden {
return true
}
if statusCode/100 == 5 {
return true
}
if statusCode == http.StatusBadRequest {
return false
}
if statusCode/100 == 2 {
return false
}
return true
return false
}
func RelayNotImplemented(c *gin.Context) {
+115 -36
View File
@@ -8,12 +8,33 @@ import (
"time"
"github.com/gin-gonic/gin"
json "github.com/json-iterator/go"
"github.com/labring/sealos/service/aiproxy/common/network"
"github.com/labring/sealos/service/aiproxy/common/random"
"github.com/labring/sealos/service/aiproxy/middleware"
"github.com/labring/sealos/service/aiproxy/model"
)
type TokenResponse struct {
*model.Token
AccessedAt time.Time `json:"accessed_at"`
}
func (t *TokenResponse) MarshalJSON() ([]byte, error) {
type Alias TokenResponse
return json.Marshal(&struct {
*Alias
CreatedAt int64 `json:"created_at"`
ExpiredAt int64 `json:"expired_at"`
AccessedAt int64 `json:"accessed_at"`
}{
Alias: (*Alias)(t),
CreatedAt: t.CreatedAt.UnixMilli(),
ExpiredAt: t.ExpiredAt.UnixMilli(),
AccessedAt: t.AccessedAt.UnixMilli(),
})
}
func GetTokens(c *gin.Context) {
p, _ := strconv.Atoi(c.Query("p"))
p--
@@ -29,18 +50,31 @@ func GetTokens(c *gin.Context) {
group := c.Query("group")
order := c.Query("order")
status, _ := strconv.Atoi(c.Query("status"))
tokens, total, err := model.GetTokens(p*perPage, perPage, order, group, status)
tokens, total, err := model.GetTokens(group, p*perPage, perPage, order, status)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
tokenResponses := make([]*TokenResponse, len(tokens))
for i, token := range tokens {
lastRequestAt, _ := model.GetTokenLastRequestTime(token.ID)
tokenResponses[i] = &TokenResponse{
Token: token,
AccessedAt: lastRequestAt,
}
}
middleware.SuccessResponse(c, gin.H{
"tokens": tokens,
"tokens": tokenResponses,
"total": total,
})
}
func GetGroupTokens(c *gin.Context) {
group := c.Param("group")
if group == "" {
middleware.ErrorResponse(c, http.StatusOK, "group is required")
return
}
p, _ := strconv.Atoi(c.Query("p"))
p--
if p < 0 {
@@ -52,16 +86,23 @@ func GetGroupTokens(c *gin.Context) {
} else if perPage > 100 {
perPage = 100
}
group := c.Param("group")
order := c.Query("order")
status, _ := strconv.Atoi(c.Query("status"))
tokens, total, err := model.GetGroupTokens(group, p*perPage, perPage, order, status)
tokens, total, err := model.GetTokens(group, p*perPage, perPage, order, status)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
tokenResponses := make([]*TokenResponse, len(tokens))
for i, token := range tokens {
lastRequestAt, _ := model.GetTokenLastRequestTime(token.ID)
tokenResponses[i] = &TokenResponse{
Token: token,
AccessedAt: lastRequestAt,
}
}
middleware.SuccessResponse(c, gin.H{
"tokens": tokens,
"tokens": tokenResponses,
"total": total,
})
}
@@ -84,18 +125,31 @@ func SearchTokens(c *gin.Context) {
key := c.Query("key")
status, _ := strconv.Atoi(c.Query("status"))
group := c.Query("group")
tokens, total, err := model.SearchTokens(keyword, p*perPage, perPage, order, status, name, key, group)
tokens, total, err := model.SearchTokens(group, keyword, p*perPage, perPage, order, status, name, key)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
tokenResponses := make([]*TokenResponse, len(tokens))
for i, token := range tokens {
lastRequestAt, _ := model.GetTokenLastRequestTime(token.ID)
tokenResponses[i] = &TokenResponse{
Token: token,
AccessedAt: lastRequestAt,
}
}
middleware.SuccessResponse(c, gin.H{
"tokens": tokens,
"tokens": tokenResponses,
"total": total,
})
}
func SearchGroupTokens(c *gin.Context) {
group := c.Param("group")
if group == "" {
middleware.ErrorResponse(c, http.StatusOK, "group is required")
return
}
keyword := c.Query("keyword")
p, _ := strconv.Atoi(c.Query("p"))
p--
@@ -108,18 +162,25 @@ func SearchGroupTokens(c *gin.Context) {
} else if perPage > 100 {
perPage = 100
}
group := c.Param("group")
order := c.Query("order")
name := c.Query("name")
key := c.Query("key")
status, _ := strconv.Atoi(c.Query("status"))
tokens, total, err := model.SearchGroupTokens(group, keyword, p*perPage, perPage, order, status, name, key)
tokens, total, err := model.SearchTokens(group, keyword, p*perPage, perPage, order, status, name, key)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
tokenResponses := make([]*TokenResponse, len(tokens))
for i, token := range tokens {
lastRequestAt, _ := model.GetTokenLastRequestTime(token.ID)
tokenResponses[i] = &TokenResponse{
Token: token,
AccessedAt: lastRequestAt,
}
}
middleware.SuccessResponse(c, gin.H{
"tokens": tokens,
"tokens": tokenResponses,
"total": total,
})
}
@@ -135,22 +196,36 @@ func GetToken(c *gin.Context) {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
middleware.SuccessResponse(c, token)
lastRequestAt, _ := model.GetTokenLastRequestTime(id)
tokenResponse := &TokenResponse{
Token: token,
AccessedAt: lastRequestAt,
}
middleware.SuccessResponse(c, tokenResponse)
}
func GetGroupToken(c *gin.Context) {
group := c.Param("group")
if group == "" {
middleware.ErrorResponse(c, http.StatusOK, "group is required")
return
}
id, err := strconv.Atoi(c.Param("id"))
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
group := c.Param("group")
token, err := model.GetGroupTokenByID(group, id)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
middleware.SuccessResponse(c, token)
lastRequestAt, _ := model.GetTokenLastRequestTime(id)
tokenResponse := &TokenResponse{
Token: token,
AccessedAt: lastRequestAt,
}
middleware.SuccessResponse(c, tokenResponse)
}
func validateToken(token AddTokenRequest) error {
@@ -212,7 +287,9 @@ func AddToken(c *gin.Context) {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
middleware.SuccessResponse(c, cleanToken)
middleware.SuccessResponse(c, &TokenResponse{
Token: cleanToken,
})
}
func DeleteToken(c *gin.Context) {
@@ -311,7 +388,9 @@ func UpdateToken(c *gin.Context) {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
middleware.SuccessResponse(c, cleanToken)
middleware.SuccessResponse(c, &TokenResponse{
Token: cleanToken,
})
}
func UpdateGroupToken(c *gin.Context) {
@@ -351,7 +430,9 @@ func UpdateGroupToken(c *gin.Context) {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
middleware.SuccessResponse(c, cleanToken)
middleware.SuccessResponse(c, &TokenResponse{
Token: cleanToken,
})
}
type UpdateTokenStatusRequest struct {
@@ -375,20 +456,14 @@ func UpdateTokenStatus(c *gin.Context) {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
if token.Status == model.TokenStatusEnabled {
if cleanToken.Status == model.TokenStatusExpired && !cleanToken.ExpiredAt.IsZero() && cleanToken.ExpiredAt.Before(time.Now()) {
middleware.ErrorResponse(c, http.StatusOK, "token expired, please update token expired time or set to never expire")
return
}
if cleanToken.Status == model.TokenStatusExhausted && cleanToken.Quota > 0 && cleanToken.UsedAmount >= cleanToken.Quota {
middleware.ErrorResponse(c, http.StatusOK, "token quota exhausted, please update token quota or set to unlimited quota")
return
}
if cleanToken.Status == model.TokenStatusExhausted && cleanToken.Quota > 0 && cleanToken.UsedAmount >= cleanToken.Quota {
middleware.ErrorResponse(c, http.StatusOK, "token quota exhausted, please update token quota or set to unlimited quota")
if err := validateTokenStatus(cleanToken); err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
}
err = model.UpdateTokenStatus(id, token.Status)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
@@ -419,20 +494,14 @@ func UpdateGroupTokenStatus(c *gin.Context) {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
if token.Status == model.TokenStatusEnabled {
if cleanToken.Status == model.TokenStatusExpired && !cleanToken.ExpiredAt.IsZero() && cleanToken.ExpiredAt.Before(time.Now()) {
middleware.ErrorResponse(c, http.StatusOK, "token expired, please update token expired time or set to never expire")
return
}
if cleanToken.Status == model.TokenStatusExhausted && cleanToken.Quota > 0 && cleanToken.UsedAmount >= cleanToken.Quota {
middleware.ErrorResponse(c, http.StatusOK, "token quota exhausted, please update token quota or set to unlimited quota")
return
}
if cleanToken.Status == model.TokenStatusExhausted && cleanToken.Quota > 0 && cleanToken.UsedAmount >= cleanToken.Quota {
middleware.ErrorResponse(c, http.StatusOK, "token quota exhausted, please update token quota or set to unlimited quota")
if err := validateTokenStatus(cleanToken); err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
return
}
}
err = model.UpdateGroupTokenStatus(group, id, token.Status)
if err != nil {
middleware.ErrorResponse(c, http.StatusOK, err.Error())
@@ -441,6 +510,16 @@ func UpdateGroupTokenStatus(c *gin.Context) {
middleware.SuccessResponse(c, nil)
}
func validateTokenStatus(token *model.Token) error {
if token.Status == model.TokenStatusExpired && !token.ExpiredAt.IsZero() && token.ExpiredAt.Before(time.Now()) {
return errors.New("token expired, please update token expired time or set to never expire")
}
if token.Status == model.TokenStatusExhausted && token.Quota > 0 && token.UsedAmount >= token.Quota {
return errors.New("token quota exhausted, please update token quota or set to unlimited quota")
}
return nil
}
type UpdateTokenNameRequest struct {
Name string `json:"name"`
}
+3
View File
@@ -13,4 +13,7 @@ ENV SQL_DSN="<sql-placeholder>"
ENV LOG_SQL_DSN="<sql-log-placeholder>"
ENV REDIS_CONN_STRING="<redis-placeholder>"
ENV BALANCE_SEALOS_CHECK_REAL_NAME_ENABLE="false"
ENV BALANCE_SEALOS_NO_REAL_NAME_USED_AMOUNT_LIMIT="1"
CMD ["bash scripts/init.sh"]
@@ -10,3 +10,5 @@ data:
SQL_DSN: "{{ .SQL_DSN }}"
LOG_SQL_DSN: "{{ .LOG_SQL_DSN }}"
REDIS_CONN_STRING: "{{ .REDIS_CONN_STRING }}"
BALANCE_SEALOS_CHECK_REAL_NAME_ENABLE: "{{ .BALANCE_SEALOS_CHECK_REAL_NAME_ENABLE }}"
BALANCE_SEALOS_NO_REAL_NAME_USED_AMOUNT_LIMIT: "{{ .BALANCE_SEALOS_NO_REAL_NAME_USED_AMOUNT_LIMIT }}"
+51 -52
View File
@@ -5,12 +5,12 @@ go 1.22.7
replace github.com/labring/sealos/service/aiproxy => ../aiproxy
require (
cloud.google.com/go/iam v1.3.0
github.com/aws/aws-sdk-go-v2 v1.32.6
github.com/aws/aws-sdk-go-v2/credentials v1.17.47
github.com/aws/aws-sdk-go-v2/service/bedrockruntime v1.23.0
github.com/gin-contrib/cors v1.7.2
github.com/gin-contrib/gzip v1.0.1
cloud.google.com/go/iam v1.3.1
github.com/aws/aws-sdk-go-v2 v1.36.1
github.com/aws/aws-sdk-go-v2/credentials v1.17.59
github.com/aws/aws-sdk-go-v2/service/bedrockruntime v1.24.4
github.com/gin-contrib/cors v1.7.3
github.com/gin-contrib/gzip v1.2.2
github.com/gin-gonic/gin v1.10.0
github.com/glebarez/sqlite v1.11.0
github.com/golang-jwt/jwt/v5 v5.2.1
@@ -28,57 +28,57 @@ require (
github.com/shopspring/decimal v1.4.0
github.com/sirupsen/logrus v1.9.3
github.com/smartystreets/goconvey v1.8.1
github.com/stretchr/testify v1.9.0
golang.org/x/image v0.23.0
google.golang.org/api v0.210.0
github.com/srwiley/oksvg v0.0.0-20221011165216-be6e8873101c
github.com/srwiley/rasterx v0.0.0-20220730225603-2ab79fcdd4ef
github.com/stretchr/testify v1.10.0
golang.org/x/image v0.24.0
golang.org/x/sync v0.11.0
google.golang.org/api v0.220.0
gorm.io/driver/mysql v1.5.7
gorm.io/driver/postgres v1.5.11
gorm.io/gorm v1.25.12
)
require (
cloud.google.com/go/auth v0.12.0 // indirect
cloud.google.com/go/auth/oauth2adapt v0.2.6 // indirect
cloud.google.com/go/compute/metadata v0.5.2 // indirect
cloud.google.com/go/auth v0.14.1 // indirect
cloud.google.com/go/auth/oauth2adapt v0.2.7 // indirect
cloud.google.com/go/compute/metadata v0.6.0 // indirect
filippo.io/edwards25519 v1.1.0 // indirect
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.6.7 // indirect
github.com/aws/aws-sdk-go-v2/internal/configsources v1.3.25 // indirect
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.6.25 // indirect
github.com/aws/smithy-go v1.22.1 // indirect
github.com/bytedance/sonic v1.12.5 // indirect
github.com/bytedance/sonic/loader v0.2.1 // indirect
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.6.8 // indirect
github.com/aws/aws-sdk-go-v2/internal/configsources v1.3.32 // indirect
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.6.32 // indirect
github.com/aws/smithy-go v1.22.2 // indirect
github.com/bytedance/sonic v1.12.8 // indirect
github.com/bytedance/sonic/loader v0.2.3 // indirect
github.com/cespare/xxhash/v2 v2.3.0 // indirect
github.com/cloudwego/base64x v0.1.4 // indirect
github.com/cloudwego/iasm v0.2.0 // indirect
github.com/cloudwego/base64x v0.1.5 // indirect
github.com/davecgh/go-spew v1.1.1 // indirect
github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect
github.com/dlclark/regexp2 v1.11.4 // indirect
github.com/dustin/go-humanize v1.0.1 // indirect
github.com/felixge/httpsnoop v1.0.4 // indirect
github.com/gabriel-vasile/mimetype v1.4.7 // indirect
github.com/gin-contrib/sse v0.1.0 // indirect
github.com/gabriel-vasile/mimetype v1.4.8 // indirect
github.com/gin-contrib/sse v1.0.0 // indirect
github.com/glebarez/go-sqlite v1.22.0 // indirect
github.com/go-logr/logr v1.4.2 // indirect
github.com/go-logr/stdr v1.2.2 // indirect
github.com/go-playground/locales v0.14.1 // indirect
github.com/go-playground/universal-translator v0.18.1 // indirect
github.com/go-playground/validator/v10 v10.23.0 // indirect
github.com/go-playground/validator/v10 v10.24.0 // indirect
github.com/go-sql-driver/mysql v1.8.1 // indirect
github.com/goccy/go-json v0.10.3 // indirect
github.com/golang/groupcache v0.0.0-20241129210726-2c02b8208cf8 // indirect
github.com/google/s2a-go v0.1.8 // indirect
github.com/goccy/go-json v0.10.5 // indirect
github.com/google/s2a-go v0.1.9 // indirect
github.com/googleapis/enterprise-certificate-proxy v0.3.4 // indirect
github.com/googleapis/gax-go/v2 v2.14.0 // indirect
github.com/googleapis/gax-go/v2 v2.14.1 // indirect
github.com/gopherjs/gopherjs v1.17.2 // indirect
github.com/jackc/pgpassfile v1.0.0 // indirect
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect
github.com/jackc/pgx/v5 v5.7.1 // indirect
github.com/jackc/pgx/v5 v5.7.2 // indirect
github.com/jackc/puddle/v2 v2.2.2 // indirect
github.com/jinzhu/inflection v1.0.0 // indirect
github.com/jinzhu/now v1.1.5 // indirect
github.com/jtolds/gls v4.20.0+incompatible // indirect
github.com/klauspost/cpuid/v2 v2.2.9 // indirect
github.com/kr/text v0.2.0 // indirect
github.com/leodido/go-urn v1.4.0 // indirect
github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect
github.com/modern-go/reflect2 v1.0.2 // indirect
@@ -89,28 +89,27 @@ require (
github.com/smarty/assertions v1.15.0 // indirect
github.com/twitchyliquid64/golang-asm v0.15.1 // indirect
github.com/ugorji/go/codec v1.2.12 // indirect
go.opencensus.io v0.24.0 // indirect
go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.57.0 // indirect
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.57.0 // indirect
go.opentelemetry.io/otel v1.32.0 // indirect
go.opentelemetry.io/otel/metric v1.32.0 // indirect
go.opentelemetry.io/otel/trace v1.32.0 // indirect
golang.org/x/arch v0.12.0 // indirect
golang.org/x/crypto v0.30.0 // indirect
golang.org/x/exp v0.0.0-20241204233417-43b7b7cde48d // indirect
golang.org/x/net v0.32.0 // indirect
golang.org/x/oauth2 v0.24.0 // indirect
golang.org/x/sync v0.10.0 // indirect
golang.org/x/sys v0.28.0 // indirect
golang.org/x/text v0.21.0 // indirect
golang.org/x/time v0.8.0 // indirect
google.golang.org/genproto/googleapis/api v0.0.0-20241209162323-e6fa225c2576 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20241209162323-e6fa225c2576 // indirect
google.golang.org/grpc v1.68.1 // indirect
google.golang.org/protobuf v1.35.2 // indirect
go.opentelemetry.io/auto/sdk v1.1.0 // indirect
go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.59.0 // indirect
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.59.0 // indirect
go.opentelemetry.io/otel v1.34.0 // indirect
go.opentelemetry.io/otel/metric v1.34.0 // indirect
go.opentelemetry.io/otel/trace v1.34.0 // indirect
golang.org/x/arch v0.14.0 // indirect
golang.org/x/crypto v0.33.0 // indirect
golang.org/x/exp v0.0.0-20250210185358-939b2ce775ac // indirect
golang.org/x/net v0.35.0 // indirect
golang.org/x/oauth2 v0.26.0 // indirect
golang.org/x/sys v0.30.0 // indirect
golang.org/x/text v0.22.0 // indirect
golang.org/x/time v0.10.0 // indirect
google.golang.org/genproto/googleapis/api v0.0.0-20250207221924-e9438ea467c6 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20250207221924-e9438ea467c6 // indirect
google.golang.org/grpc v1.70.0 // indirect
google.golang.org/protobuf v1.36.5 // indirect
gopkg.in/yaml.v3 v3.0.1 // indirect
modernc.org/libc v1.61.4 // indirect
modernc.org/mathutil v1.6.0 // indirect
modernc.org/memory v1.8.0 // indirect
modernc.org/sqlite v1.34.2 // indirect
modernc.org/libc v1.61.12 // indirect
modernc.org/mathutil v1.7.1 // indirect
modernc.org/memory v1.8.2 // indirect
modernc.org/sqlite v1.34.5 // indirect
)
+128 -200
View File
@@ -1,48 +1,41 @@
cloud.google.com/go v0.26.0/go.mod h1:aQUYkXzVsufM+DwF1aE+0xfcU+56JwCaLick0ClmMTw=
cloud.google.com/go/auth v0.12.0 h1:ARAD8r0lkiHw2go7kEnmviF6TOYhzLM+yDGcDt9mP68=
cloud.google.com/go/auth v0.12.0/go.mod h1:xxA5AqpDrvS+Gkmo9RqrGGRh6WSNKKOXhY3zNOr38tI=
cloud.google.com/go/auth/oauth2adapt v0.2.6 h1:V6a6XDu2lTwPZWOawrAa9HUK+DB2zfJyTuciBG5hFkU=
cloud.google.com/go/auth/oauth2adapt v0.2.6/go.mod h1:AlmsELtlEBnaNTL7jCj8VQFLy6mbZv0s4Q7NGBeQ5E8=
cloud.google.com/go/compute/metadata v0.5.2 h1:UxK4uu/Tn+I3p2dYWTfiX4wva7aYlKixAHn3fyqngqo=
cloud.google.com/go/compute/metadata v0.5.2/go.mod h1:C66sj2AluDcIqakBq/M8lw8/ybHgOZqin2obFxa/E5k=
cloud.google.com/go/iam v1.3.0 h1:4Wo2qTaGKFtajbLpF6I4mywg900u3TLlHDb6mriLDPU=
cloud.google.com/go/iam v1.3.0/go.mod h1:0Ys8ccaZHdI1dEUilwzqng/6ps2YB6vRsjIe00/+6JY=
cloud.google.com/go/auth v0.14.1 h1:AwoJbzUdxA/whv1qj3TLKwh3XX5sikny2fc40wUl+h0=
cloud.google.com/go/auth v0.14.1/go.mod h1:4JHUxlGXisL0AW8kXPtUF6ztuOksyfUQNFjfsOCXkPM=
cloud.google.com/go/auth/oauth2adapt v0.2.7 h1:/Lc7xODdqcEw8IrZ9SvwnlLX6j9FHQM74z6cBk9Rw6M=
cloud.google.com/go/auth/oauth2adapt v0.2.7/go.mod h1:NTbTTzfvPl1Y3V1nPpOgl2w6d/FjO7NNUQaWSox6ZMc=
cloud.google.com/go/compute/metadata v0.6.0 h1:A6hENjEsCDtC1k8byVsgwvVcioamEHvZ4j01OwKxG9I=
cloud.google.com/go/compute/metadata v0.6.0/go.mod h1:FjyFAW1MW0C203CEOMDTu3Dk1FlqW3Rga40jzHL4hfg=
cloud.google.com/go/iam v1.3.1 h1:KFf8SaT71yYq+sQtRISn90Gyhyf4X8RGgeAVC8XGf3E=
cloud.google.com/go/iam v1.3.1/go.mod h1:3wMtuyT4NcbnYNPLMBzYRFiEfjKfJlLVLrisE7bwm34=
filippo.io/edwards25519 v1.1.0 h1:FNf4tywRC1HmFuKW5xopWpigGjJKiJSV0Cqo0cJWDaA=
filippo.io/edwards25519 v1.1.0/go.mod h1:BxyFTGdWcka3PhytdK4V28tE5sGfRvvvRV7EaN4VDT4=
github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU=
github.com/aws/aws-sdk-go-v2 v1.32.6 h1:7BokKRgRPuGmKkFMhEg/jSul+tB9VvXhcViILtfG8b4=
github.com/aws/aws-sdk-go-v2 v1.32.6/go.mod h1:P5WJBrYqqbWVaOxgH0X/FYYD47/nooaPOZPlQdmiN2U=
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.6.7 h1:lL7IfaFzngfx0ZwUGOZdsFFnQ5uLvR0hWqqhyE7Q9M8=
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.6.7/go.mod h1:QraP0UcVlQJsmHfioCrveWOC1nbiWUl3ej08h4mXWoc=
github.com/aws/aws-sdk-go-v2/credentials v1.17.47 h1:48bA+3/fCdi2yAwVt+3COvmatZ6jUDNkDTIsqDiMUdw=
github.com/aws/aws-sdk-go-v2/credentials v1.17.47/go.mod h1:+KdckOejLW3Ks3b0E3b5rHsr2f9yuORBum0WPnE5o5w=
github.com/aws/aws-sdk-go-v2/internal/configsources v1.3.25 h1:s/fF4+yDQDoElYhfIVvSNyeCydfbuTKzhxSXDXCPasU=
github.com/aws/aws-sdk-go-v2/internal/configsources v1.3.25/go.mod h1:IgPfDv5jqFIzQSNbUEMoitNooSMXjRSDkhXv8jiROvU=
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.6.25 h1:ZntTCl5EsYnhN/IygQEUugpdwbhdkom9uHcbCftiGgA=
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.6.25/go.mod h1:DBdPrgeocww+CSl1C8cEV8PN1mHMBhuCDLpXezyvWkE=
github.com/aws/aws-sdk-go-v2/service/bedrockruntime v1.23.0 h1:mfV5tcLXeRLbiyI4EHoHWH1sIU7JvbfXVvymUCIgZEo=
github.com/aws/aws-sdk-go-v2/service/bedrockruntime v1.23.0/go.mod h1:YSSgYnasDKm5OjU3bOPkaz+2PFO6WjEQGIA6KQNsR3Q=
github.com/aws/smithy-go v1.22.1 h1:/HPHZQ0g7f4eUeK6HKglFz8uwVfZKgoI25rb/J+dnro=
github.com/aws/smithy-go v1.22.1/go.mod h1:irrKGvNn1InZwb2d7fkIRNucdfwR8R+Ts3wxYa/cJHg=
github.com/aws/aws-sdk-go-v2 v1.36.1 h1:iTDl5U6oAhkNPba0e1t1hrwAo02ZMqbrGq4k5JBWM5E=
github.com/aws/aws-sdk-go-v2 v1.36.1/go.mod h1:5PMILGVKiW32oDzjj6RU52yrNrDPUHcbZQYr1sM7qmM=
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.6.8 h1:zAxi9p3wsZMIaVCdoiQp2uZ9k1LsZvmAnoTBeZPXom0=
github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.6.8/go.mod h1:3XkePX5dSaxveLAYY7nsbsZZrKxCyEuE5pM4ziFxyGg=
github.com/aws/aws-sdk-go-v2/credentials v1.17.59 h1:9btwmrt//Q6JcSdgJOLI98sdr5p7tssS9yAsGe8aKP4=
github.com/aws/aws-sdk-go-v2/credentials v1.17.59/go.mod h1:NM8fM6ovI3zak23UISdWidyZuI1ghNe2xjzUZAyT+08=
github.com/aws/aws-sdk-go-v2/internal/configsources v1.3.32 h1:BjUcr3X3K0wZPGFg2bxOWW3VPN8rkE3/61zhP+IHviA=
github.com/aws/aws-sdk-go-v2/internal/configsources v1.3.32/go.mod h1:80+OGC/bgzzFFTUmcuwD0lb4YutwQeKLFpmt6hoWapU=
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.6.32 h1:m1GeXHVMJsRsUAqG6HjZWx9dj7F5TR+cF1bjyfYyBd4=
github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.6.32/go.mod h1:IitoQxGfaKdVLNg0hD8/DXmAqNy0H4K2H2Sf91ti8sI=
github.com/aws/aws-sdk-go-v2/service/bedrockruntime v1.24.4 h1:NYHDOBe0ZIeQfaPSPRaQym2NePzA+QYM3O/Oh4IznKg=
github.com/aws/aws-sdk-go-v2/service/bedrockruntime v1.24.4/go.mod h1:AD+JAcEr9fNzFcfKs3CINKBdWGFK7R+/uZ+VdJRhK2U=
github.com/aws/smithy-go v1.22.2 h1:6D9hW43xKFrRx/tXXfAlIZc4JI+yQe6snnWcQyxSyLQ=
github.com/aws/smithy-go v1.22.2/go.mod h1:irrKGvNn1InZwb2d7fkIRNucdfwR8R+Ts3wxYa/cJHg=
github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs=
github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c=
github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA=
github.com/bsm/gomega v1.27.10/go.mod h1:JyEr/xRbxbtgWNi8tIEVPUYZ5Dzef52k01W3YH0H+O0=
github.com/bytedance/sonic v1.12.5 h1:hoZxY8uW+mT+OpkcUWw4k0fDINtOcVavEsGfzwzFU/w=
github.com/bytedance/sonic v1.12.5/go.mod h1:B8Gt/XvtZ3Fqj+iSKMypzymZxw/FVwgIGKzMzT9r/rk=
github.com/bytedance/sonic v1.12.8 h1:4xYRVRlXIgvSZ4e8iVTlMF5szgpXd4AfvuWgA8I8lgs=
github.com/bytedance/sonic v1.12.8/go.mod h1:uVvFidNmlt9+wa31S1urfwwthTWteBgG0hWuoKAXTx8=
github.com/bytedance/sonic/loader v0.1.1/go.mod h1:ncP89zfokxS5LZrJxl5z0UJcsk4M4yY2JpfqGeCtNLU=
github.com/bytedance/sonic/loader v0.2.1 h1:1GgorWTqf12TA8mma4DDSbaQigE2wOgQo7iCjjJv3+E=
github.com/bytedance/sonic/loader v0.2.1/go.mod h1:ncP89zfokxS5LZrJxl5z0UJcsk4M4yY2JpfqGeCtNLU=
github.com/census-instrumentation/opencensus-proto v0.2.1/go.mod h1:f6KPmirojxKA12rnyqOA5BBL4O983OfeGPqjHWSTneU=
github.com/bytedance/sonic/loader v0.2.3 h1:yctD0Q3v2NOGfSWPLPvG2ggA2kV6TS6s4wioyEqssH0=
github.com/bytedance/sonic/loader v0.2.3/go.mod h1:N8A3vUdtUebEY2/VQC0MyhYeKUFosQU6FxH2JmUe6VI=
github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
github.com/client9/misspell v0.3.4/go.mod h1:qj6jICC3Q7zFZvVWo7KLAzC3yx5G7kyvSDkc90ppPyw=
github.com/cloudwego/base64x v0.1.4 h1:jwCgWpFanWmN8xoIUHa2rtzmkd5J2plF/dnLS6Xd/0Y=
github.com/cloudwego/base64x v0.1.4/go.mod h1:0zlkT4Wn5C6NdauXdJRhSKRlJvmclQ1hhJgA0rcu/8w=
github.com/cloudwego/iasm v0.2.0 h1:1KNIy1I1H9hNNFEEH3DVnI4UujN+1zjpuk6gwHLTssg=
github.com/cloudwego/base64x v0.1.5 h1:XPciSp1xaq2VCSt6lF0phncD4koWyULpl5bUxbfCyP4=
github.com/cloudwego/base64x v0.1.5/go.mod h1:0zlkT4Wn5C6NdauXdJRhSKRlJvmclQ1hhJgA0rcu/8w=
github.com/cloudwego/iasm v0.2.0/go.mod h1:8rXZaNYT2n95jn+zTI1sDr+IgcD2GVs0nlbbQPiEFhY=
github.com/cncf/udpa/go v0.0.0-20191209042840-269d4d468f6f/go.mod h1:M8M6+tZqaGXZJjfX53e64911xZQV5JYwmTeXPW+k8Sc=
github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ33E=
github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
@@ -52,20 +45,16 @@ github.com/dlclark/regexp2 v1.11.4 h1:rPYF9/LECdNymJufQKmri9gV604RvvABwgOA8un7yA
github.com/dlclark/regexp2 v1.11.4/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cnmRbL6yW8=
github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkpeCY=
github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto=
github.com/envoyproxy/go-control-plane v0.9.0/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4=
github.com/envoyproxy/go-control-plane v0.9.1-0.20191026205805-5f8ba28d4473/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4=
github.com/envoyproxy/go-control-plane v0.9.4/go.mod h1:6rpuAdCZL397s3pYoYcLgu1mIlRU8Am5FuJP05cCM98=
github.com/envoyproxy/protoc-gen-validate v0.1.0/go.mod h1:iSmxcyjqTsJpI2R4NaDN7+kN2VEUnK/pcBlmesArF7c=
github.com/felixge/httpsnoop v1.0.4 h1:NFTV2Zj1bL4mc9sqWACXbQFVBBg2W3GPvqp8/ESS2Wg=
github.com/felixge/httpsnoop v1.0.4/go.mod h1:m8KPJKqk1gH5J9DgRY2ASl2lWCfGKXixSwevea8zH2U=
github.com/gabriel-vasile/mimetype v1.4.7 h1:SKFKl7kD0RiPdbht0s7hFtjl489WcQ1VyPW8ZzUMYCA=
github.com/gabriel-vasile/mimetype v1.4.7/go.mod h1:GDlAgAyIRT27BhFl53XNAFtfjzOkLaF35JdEG0P7LtU=
github.com/gin-contrib/cors v1.7.2 h1:oLDHxdg8W/XDoN/8zamqk/Drgt4oVZDvaV0YmvVICQw=
github.com/gin-contrib/cors v1.7.2/go.mod h1:SUJVARKgQ40dmrzgXEVxj2m7Ig1v1qIboQkPDTQ9t2E=
github.com/gin-contrib/gzip v1.0.1 h1:HQ8ENHODeLY7a4g1Au/46Z92bdGFl74OhxcZble9WJE=
github.com/gin-contrib/gzip v1.0.1/go.mod h1:njt428fdUNRvjuJf16tZMYZ2Yl+WQB53X5wmhDwXvC4=
github.com/gin-contrib/sse v0.1.0 h1:Y/yl/+YNO8GZSjAhjMsSuLt29uWRFHdHYUb5lYOV9qE=
github.com/gin-contrib/sse v0.1.0/go.mod h1:RHrZQHXnP2xjPF+u1gW/2HnVO7nvIa9PG3Gm+fLHvGI=
github.com/gabriel-vasile/mimetype v1.4.8 h1:FfZ3gj38NjllZIeJAmMhr+qKL8Wu+nOoI3GqacKw1NM=
github.com/gabriel-vasile/mimetype v1.4.8/go.mod h1:ByKUIKGjh1ODkGM1asKUbQZOLGrPjydw3hYPU2YU9t8=
github.com/gin-contrib/cors v1.7.3 h1:hV+a5xp8hwJoTw7OY+a70FsL8JkVVFTXw9EcfrYUdns=
github.com/gin-contrib/cors v1.7.3/go.mod h1:M3bcKZhxzsvI+rlRSkkxHyljJt1ESd93COUvemZ79j4=
github.com/gin-contrib/gzip v1.2.2 h1:iUU/EYCM8ENfkjmZaVrxbjF/ZC267Iqv5S0MMCMEliI=
github.com/gin-contrib/gzip v1.2.2/go.mod h1:C1a5cacjlDsS20cKnHlZRCPUu57D3qH6B2pV0rl+Y/s=
github.com/gin-contrib/sse v1.0.0 h1:y3bT1mUWUxDpW4JLQg/HnTqV4rozuW4tC9eFKTxYI9E=
github.com/gin-contrib/sse v1.0.0/go.mod h1:zNuFdwarAygJBht0NTKiSi3jRf6RbqeILZ9Sp6Slhe0=
github.com/gin-gonic/gin v1.10.0 h1:nTuyha1TYqgedzytsKYqna+DfLos46nTv2ygFy86HFU=
github.com/gin-gonic/gin v1.10.0/go.mod h1:4PMNQiOhvDRa013RKVbsiNwoyezlm2rm0uX/T7kzp5Y=
github.com/glebarez/go-sqlite v1.22.0 h1:uAcMJhaA6r3LHMTFgP0SifzgXg46yJkgxqyuyec+ruQ=
@@ -83,51 +72,30 @@ github.com/go-playground/locales v0.14.1 h1:EWaQ/wswjilfKLTECiXz7Rh+3BjFhfDFKv/o
github.com/go-playground/locales v0.14.1/go.mod h1:hxrqLVvrK65+Rwrd5Fc6F2O76J/NuW9t0sjnWqG1slY=
github.com/go-playground/universal-translator v0.18.1 h1:Bcnm0ZwsGyWbCzImXv+pAJnYK9S473LQFuzCbDbfSFY=
github.com/go-playground/universal-translator v0.18.1/go.mod h1:xekY+UJKNuX9WP91TpwSH2VMlDf28Uj24BCp08ZFTUY=
github.com/go-playground/validator/v10 v10.23.0 h1:/PwmTwZhS0dPkav3cdK9kV1FsAmrL8sThn8IHr/sO+o=
github.com/go-playground/validator/v10 v10.23.0/go.mod h1:dbuPbCMFw/DrkbEynArYaCwl3amGuJotoKCe95atGMM=
github.com/go-playground/validator/v10 v10.24.0 h1:KHQckvo8G6hlWnrPX4NJJ+aBfWNAE/HH+qdL2cBpCmg=
github.com/go-playground/validator/v10 v10.24.0/go.mod h1:GGzBIJMuE98Ic/kJsBXbz1x/7cByt++cQ+YOuDM5wus=
github.com/go-sql-driver/mysql v1.7.0/go.mod h1:OXbVy3sEdcQ2Doequ6Z5BW6fXNQTmx+9S1MCJN5yJMI=
github.com/go-sql-driver/mysql v1.8.1 h1:LedoTUt/eveggdHS9qUFC1EFSa8bU2+1pZjSRpvNJ1Y=
github.com/go-sql-driver/mysql v1.8.1/go.mod h1:wEBSXgmK//2ZFJyE+qWnIsVGmvmEKlqwuVSjsCm7DZg=
github.com/goccy/go-json v0.10.3 h1:KZ5WoDbxAIgm2HNbYckL0se1fHD6rz5j4ywS6ebzDqA=
github.com/goccy/go-json v0.10.3/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M=
github.com/goccy/go-json v0.10.5 h1:Fq85nIqj+gXn/S5ahsiTlK3TmC85qgirsdTP/+DeaC4=
github.com/goccy/go-json v0.10.5/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M=
github.com/golang-jwt/jwt/v5 v5.2.1 h1:OuVbFODueb089Lh128TAcimifWaLhJwVflnrgM17wHk=
github.com/golang-jwt/jwt/v5 v5.2.1/go.mod h1:pqrtFR0X4osieyHYxtmOUWsAWrfe1Q5UVIyoH402zdk=
github.com/golang/glog v0.0.0-20160126235308-23def4e6c14b/go.mod h1:SBH7ygxi8pfUlaOkMMuAQtPIUF8ecWP5IEl/CR7VP2Q=
github.com/golang/groupcache v0.0.0-20200121045136-8c9f03a8e57e/go.mod h1:cIg4eruTrX1D+g88fzRXU5OdNfaM+9IcxsU14FzY7Hc=
github.com/golang/groupcache v0.0.0-20241129210726-2c02b8208cf8 h1:f+oWsMOmNPc8JmEHVZIycC7hBoQxHH9pNKQORJNozsQ=
github.com/golang/groupcache v0.0.0-20241129210726-2c02b8208cf8/go.mod h1:wcDNUvekVysuuOpQKo3191zZyTpiI6se1N1ULghS0sw=
github.com/golang/mock v1.1.1/go.mod h1:oTYuIxOrZwtPieC+H1uAHpcLFnEyAGVDL/k47Jfbm0A=
github.com/golang/protobuf v1.2.0/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U=
github.com/golang/protobuf v1.3.2/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U=
github.com/golang/protobuf v1.4.0-rc.1/go.mod h1:ceaxUfeHdC40wWswd/P6IGgMaK3YpKi5j83Wpe3EHw8=
github.com/golang/protobuf v1.4.0-rc.1.0.20200221234624-67d41d38c208/go.mod h1:xKAWHe0F5eneWXFV3EuXVDTCmh+JuBKY0li0aMyXATA=
github.com/golang/protobuf v1.4.0-rc.2/go.mod h1:LlEzMj4AhA7rCAGe4KMBDvJI+AwstrUpVNzEA03Pprs=
github.com/golang/protobuf v1.4.0-rc.4.0.20200313231945-b860323f09d0/go.mod h1:WU3c8KckQ9AFe+yFwt9sWVRKCVIyN9cPHBJSNnbL67w=
github.com/golang/protobuf v1.4.0/go.mod h1:jodUvKwWbYaEsadDk5Fwe5c77LiNKVO9IDvqG2KuDX0=
github.com/golang/protobuf v1.4.1/go.mod h1:U8fpvMrcmy5pZrNK1lt4xCsGvpyWQ/VVv6QDs8UjoX8=
github.com/golang/protobuf v1.4.3/go.mod h1:oDoupMAO8OvCJWAcko0GGGIgR6R6ocIYbsSw735rRwI=
github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek=
github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps=
github.com/google/go-cmp v0.2.0/go.mod h1:oXzfMopK8JAjlY9xF4vHSVASa0yLyX7SntLO5aqRK0M=
github.com/google/go-cmp v0.3.0/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU=
github.com/google/go-cmp v0.3.1/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU=
github.com/google/go-cmp v0.4.0/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
github.com/google/go-cmp v0.5.0/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
github.com/google/go-cmp v0.5.3/go.mod h1:v8dTdLbMG2kIc/vJvl+f65V22dbkXbowE6jgT/gNBxE=
github.com/google/go-cmp v0.6.0 h1:ofyhxvXcZhMsU5ulbFiLKl/XBFqE1GSq7atu8tAmTRI=
github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY=
github.com/google/gofuzz v1.0.0/go.mod h1:dBl0BpW6vV/+mYPU4Po3pmUjxk6FQPldtuIdl/M65Eg=
github.com/google/pprof v0.0.0-20240409012703-83162a5b38cd h1:gbpYu9NMq8jhDVbvlGkMFWCjLFlqqEZjEmObmhUy6Vo=
github.com/google/pprof v0.0.0-20240409012703-83162a5b38cd/go.mod h1:kf6iHlnVGwgKolg33glAes7Yg/8iWP8ukqeldJSO7jw=
github.com/google/s2a-go v0.1.8 h1:zZDs9gcbt9ZPLV0ndSyQk6Kacx2g/X+SKYovpnz3SMM=
github.com/google/s2a-go v0.1.8/go.mod h1:6iNWHTpQ+nfNRN5E00MSdfDwVesa8hhS32PhPO8deJA=
github.com/google/uuid v1.1.2/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/google/s2a-go v0.1.9 h1:LGD7gtMgezd8a/Xak7mEWL0PjoTQFvpRudN895yqKW0=
github.com/google/s2a-go v0.1.9/go.mod h1:YA0Ei2ZQL3acow2O62kdp9UlnvMmU7kA6Eutn0dXayM=
github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
github.com/googleapis/enterprise-certificate-proxy v0.3.4 h1:XYIDZApgAnrN1c855gTgghdIA6Stxb52D5RnLI1SLyw=
github.com/googleapis/enterprise-certificate-proxy v0.3.4/go.mod h1:YKe7cfqYXjKGpGvmSg28/fFvhNzinZQm8DGnaburhGA=
github.com/googleapis/gax-go/v2 v2.14.0 h1:f+jMrjBPl+DL9nI4IQzLUxMq7XrAqFYB7hBPqMNIe8o=
github.com/googleapis/gax-go/v2 v2.14.0/go.mod h1:lhBCnjdLrWRaPvLWhmc8IS24m9mr07qSYnHncrgo+zk=
github.com/googleapis/gax-go/v2 v2.14.1 h1:hb0FFeiPaQskmvakKu5EbCbpntQn48jyHuvrkurSS/Q=
github.com/googleapis/gax-go/v2 v2.14.1/go.mod h1:Hb/NubMaVM88SrNkvl8X/o8XWwDJEPqouaLeN2IUxoA=
github.com/gopherjs/gopherjs v1.17.2 h1:fQnZVsXk8uxXIStYb0N4bGk7jeyTalG/wsZjQ25dO0g=
github.com/gopherjs/gopherjs v1.17.2/go.mod h1:pRRIvn/QzFLrKfvEz3qUuEhtE/zLCWfreZ6J5gM2i+k=
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
@@ -136,8 +104,8 @@ github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsI
github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg=
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo=
github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM=
github.com/jackc/pgx/v5 v5.7.1 h1:x7SYsPBYDkHDksogeSmZZ5xzThcTgRz++I5E+ePFUcs=
github.com/jackc/pgx/v5 v5.7.1/go.mod h1:e7O26IywZZ+naJtWWos6i6fvWK+29etgITqrqHLfoZA=
github.com/jackc/pgx/v5 v5.7.2 h1:mLoDLV6sonKlvjIEsV56SkWNCnuNv531l94GaIzO+XI=
github.com/jackc/pgx/v5 v5.7.2/go.mod h1:ncY89UGWxg82EykZUwSpUKEfccBGGYq1xjrOpsbsfGQ=
github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo=
github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
github.com/jinzhu/copier v0.4.0 h1:w3ciUoD19shMCRargcpm0cm91ytaBhDvuRpz1ODO/U8=
@@ -156,8 +124,8 @@ github.com/klauspost/cpuid/v2 v2.0.9/go.mod h1:FInQzS24/EEf25PyTYn52gqo7WaD8xa02
github.com/klauspost/cpuid/v2 v2.2.9 h1:66ze0taIn2H33fBvCkXuv9BmCwDfafmiIVpKV9kKGuY=
github.com/klauspost/cpuid/v2 v2.2.9/go.mod h1:rqkxqrZ1EhYM9G+hXH7YdowN5R5RGN6NK4QwQ3WMXF8=
github.com/knz/go-libedit v1.10.1/go.mod h1:MZTVkCWyz0oBc7JOWP3wNAzd002ZbM/5hgShxwh4x8M=
github.com/kr/pretty v0.3.0 h1:WgNl7dwNpEZ6jJ9k1snq4pZsg7DOEN8hP9Xw0Tsjwk0=
github.com/kr/pretty v0.3.0/go.mod h1:640gp4NfQd8pI5XOwp5fnNeVWj67G7CFk/SaSQn7NBk=
github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE=
github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk=
github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY=
github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE=
github.com/leodido/go-urn v1.4.0 h1:WT9HwE9SGECu3lg4d/dIA+jxlljEa1/ffXKmRjqdmIQ=
@@ -183,13 +151,12 @@ github.com/pkoukk/tiktoken-go v0.1.7 h1:qOBHXX4PHtvIvmOtyg1EeKlwFRiMKAcoMp4Q+bLQ
github.com/pkoukk/tiktoken-go v0.1.7/go.mod h1:9NiV+i9mJKGj1rYOT+njbv+ZwA/zJxYdewGl6qVatpg=
github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM=
github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
github.com/prometheus/client_model v0.0.0-20190812154241-14fe0d1b01d4/go.mod h1:xMI15A0UPsDsEKsMN9yxemIoYk6Tm2C1GtYGdfGttqA=
github.com/redis/go-redis/v9 v9.7.0 h1:HhLSs+B6O021gwzl+locl0zEDnyNkxMtf/Z3NNBMa9E=
github.com/redis/go-redis/v9 v9.7.0/go.mod h1:f6zhXITC7JUJIlPEiBOTXxJgPLdZcA93GewI7inzyWw=
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE=
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo=
github.com/rogpeppe/go-internal v1.8.0 h1:FCbCCtXNOY3UtUuHUYaghJg4y7Fd14rXifAYUAtL9R8=
github.com/rogpeppe/go-internal v1.8.0/go.mod h1:WmiCO8CzOY8rg0OYDC4/i/2WRWAB6poM+XZ2dLUbcbE=
github.com/rogpeppe/go-internal v1.13.1 h1:KvO1DLK/DRN07sQ1LQKScxyZJuNnedQ5/wKSR38lUII=
github.com/rogpeppe/go-internal v1.13.1/go.mod h1:uMEvuHeurkdAXX61udpOXGD/AzZDWNMNyH2VO9fmH0o=
github.com/shopspring/decimal v1.4.0 h1:bxl37RwXBklmTi0C79JfXCEBD1cqqHt0bbgBAGFp81k=
github.com/shopspring/decimal v1.4.0/go.mod h1:gawqmDU56v4yIKSwfBSFip1HdCCXN8/+DMd9qYNcwME=
github.com/sirupsen/logrus v1.9.3 h1:dueUQJ1C2q9oE3F7wvmSGAaVtTmUizReu6fjN8uqzbQ=
@@ -198,115 +165,78 @@ github.com/smarty/assertions v1.15.0 h1:cR//PqUBUiQRakZWqBiFFQ9wb8emQGDb0HeGdqGB
github.com/smarty/assertions v1.15.0/go.mod h1:yABtdzeQs6l1brC900WlRNwj6ZR55d7B+E8C6HtKdec=
github.com/smartystreets/goconvey v1.8.1 h1:qGjIddxOk4grTu9JPOU31tVfq3cNdBlNa5sSznIX1xY=
github.com/smartystreets/goconvey v1.8.1/go.mod h1:+/u4qLyY6x1jReYOp7GOM2FSt8aP9CzCZL03bI28W60=
github.com/srwiley/oksvg v0.0.0-20221011165216-be6e8873101c h1:km8GpoQut05eY3GiYWEedbTT0qnSxrCjsVbb7yKY1KE=
github.com/srwiley/oksvg v0.0.0-20221011165216-be6e8873101c/go.mod h1:cNQ3dwVJtS5Hmnjxy6AgTPd0Inb3pW05ftPSX7NZO7Q=
github.com/srwiley/rasterx v0.0.0-20220730225603-2ab79fcdd4ef h1:Ch6Q+AZUxDBCVqdkI8FSpFyZDtCVBc2VmejdNrm5rRQ=
github.com/srwiley/rasterx v0.0.0-20220730225603-2ab79fcdd4ef/go.mod h1:nXTWP6+gD5+LUJ8krVhhoeHjvHTutPxMYl5SvkcnJNE=
github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME=
github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSSt89Yw=
github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo=
github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA=
github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI=
github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg=
github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU=
github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4=
github.com/stretchr/testify v1.9.0 h1:HtqpIVDClZ4nwg75+f6Lvsy/wHu+3BoSGCbBAcpTsTg=
github.com/stretchr/testify v1.9.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo=
github.com/stretchr/testify v1.10.0 h1:Xv5erBjTwe/5IxqUQTdXv5kgmIvbHo3QQyRwhJsOfJA=
github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY=
github.com/twitchyliquid64/golang-asm v0.15.1 h1:SU5vSMR7hnwNxj24w34ZyCi/FmDZTkS4MhqMhdFk5YI=
github.com/twitchyliquid64/golang-asm v0.15.1/go.mod h1:a1lVb/DtPvCB8fslRZhAngC2+aY1QWCk3Cedj/Gdt08=
github.com/ugorji/go/codec v1.2.12 h1:9LC83zGrHhuUA9l16C9AHXAqEV/2wBQ4nkvumAE65EE=
github.com/ugorji/go/codec v1.2.12/go.mod h1:UNopzCgEMSXjBc6AOMqYvWC1ktqTAfzJZUZgYf6w6lg=
go.opencensus.io v0.24.0 h1:y73uSU6J157QMP2kn2r30vwW1A2W2WFwSCGnAVxeaD0=
go.opencensus.io v0.24.0/go.mod h1:vNK8G9p7aAivkbmorf4v+7Hgx+Zs0yY+0fOtgBfjQKo=
go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.57.0 h1:qtFISDHKolvIxzSs0gIaiPUPR0Cucb0F2coHC7ZLdps=
go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.57.0/go.mod h1:Y+Pop1Q6hCOnETWTW4NROK/q1hv50hM7yDaUTjG8lp8=
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.57.0 h1:DheMAlT6POBP+gh8RUH19EOTnQIor5QE0uSRPtzCpSw=
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.57.0/go.mod h1:wZcGmeVO9nzP67aYSLDqXNWK87EZWhi7JWj1v7ZXf94=
go.opentelemetry.io/otel v1.32.0 h1:WnBN+Xjcteh0zdk01SVqV55d/m62NJLJdIyb4y/WO5U=
go.opentelemetry.io/otel v1.32.0/go.mod h1:00DCVSB0RQcnzlwyTfqtxSm+DRr9hpYrHjNGiBHVQIg=
go.opentelemetry.io/otel/metric v1.32.0 h1:xV2umtmNcThh2/a/aCP+h64Xx5wsj8qqnkYZktzNa0M=
go.opentelemetry.io/otel/metric v1.32.0/go.mod h1:jH7CIbbK6SH2V2wE16W05BHCtIDzauciCRLoc/SyMv8=
go.opentelemetry.io/otel/trace v1.32.0 h1:WIC9mYrXf8TmY/EXuULKc8hR17vE+Hjv2cssQDe03fM=
go.opentelemetry.io/otel/trace v1.32.0/go.mod h1:+i4rkvCraA+tG6AzwloGaCtkx53Fa+L+V8e9a7YvhT8=
golang.org/x/arch v0.12.0 h1:UsYJhbzPYGsT0HbEdmYcqtCv8UNGvnaL561NnIUvaKg=
golang.org/x/arch v0.12.0/go.mod h1:FEVrYAQjsQXMVJ1nsMoVVXPZg6p2JE2mx8psSWTDQys=
golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w=
golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto=
golang.org/x/crypto v0.30.0 h1:RwoQn3GkWiMkzlX562cLB7OxWvjH1L8xutO2WoJcRoY=
golang.org/x/crypto v0.30.0/go.mod h1:kDsLvtWBEx7MV9tJOj9bnXsPbxwJQ6csT/x4KIN4Ssk=
golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA=
golang.org/x/exp v0.0.0-20241204233417-43b7b7cde48d h1:0olWaB5pg3+oychR51GUVCEsGkeCU/2JxjBgIo4f3M0=
golang.org/x/exp v0.0.0-20241204233417-43b7b7cde48d/go.mod h1:qj5a5QZpwLU2NLQudwIN5koi3beDhSAlJwa67PuM98c=
golang.org/x/image v0.23.0 h1:HseQ7c2OpPKTPVzNjG5fwJsOTCiiwS4QdsYi5XU6H68=
golang.org/x/image v0.23.0/go.mod h1:wJJBTdLfCCf3tiHa1fNxpZmUI4mmoZvwMCPP0ddoNKY=
golang.org/x/lint v0.0.0-20181026193005-c67002cb31c3/go.mod h1:UVdnD1Gm6xHRNCYTkRU2/jEulfH38KcIWyp/GAMgvoE=
golang.org/x/lint v0.0.0-20190227174305-5b3e6a55c961/go.mod h1:wehouNa3lNwaWXcvxsM5YxQ5yQlVC4a0KAMCusXpPoU=
golang.org/x/lint v0.0.0-20190313153728-d0100b6bd8b3/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc=
golang.org/x/mod v0.22.0 h1:D4nJWe9zXqHOmWqj4VMOJhvzj7bEZg4wEYa759z1pH4=
golang.org/x/mod v0.22.0/go.mod h1:6SkKJ3Xj0I0BrPOZoBy3bdMptDDU9oJrpohJ3eWZ1fY=
golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20190213061140-3a22650c66bd/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4=
golang.org/x/net v0.0.0-20190311183353-d8887717615a/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg=
golang.org/x/net v0.0.0-20201110031124-69a78807bb2b/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU=
golang.org/x/net v0.32.0 h1:ZqPmj8Kzc+Y6e0+skZsuACbx+wzMgo5MQsJh9Qd6aYI=
golang.org/x/net v0.32.0/go.mod h1:CwU0IoeOlnQQWJ6ioyFrfRuomB8GKF6KbYXZVyeXNfs=
golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U=
golang.org/x/oauth2 v0.24.0 h1:KTBBxWqUa0ykRPLtV69rRto9TLXcqYkeswu48x/gvNE=
golang.org/x/oauth2 v0.24.0/go.mod h1:XYTD2NtWslqkgxebSiOHnXEap4TF09sJSc7H1sXbhtI=
golang.org/x/sync v0.0.0-20180314180146-1d60e4601c6f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20181108010431-42b317875d0f/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.0.0-20190423024810-112230192c58/go.mod h1:RxMgew5VJxzue5/jJTE5uejpjVlOe/izrB70Jof72aM=
golang.org/x/sync v0.10.0 h1:3NQrjDixjgGwUOCaF8w2+VYHv0Ve/vGYSbdkTa98gmQ=
golang.org/x/sync v0.10.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sys v0.0.0-20180830151530-49385e6e1522/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY=
golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
go.opentelemetry.io/auto/sdk v1.1.0 h1:cH53jehLUN6UFLY71z+NDOiNJqDdPRaXzTel0sJySYA=
go.opentelemetry.io/auto/sdk v1.1.0/go.mod h1:3wSPjt5PWp2RhlCcmmOial7AvC4DQqZb7a7wCow3W8A=
go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.59.0 h1:rgMkmiGfix9vFJDcDi1PK8WEQP4FLQwLDfhp5ZLpFeE=
go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc v0.59.0/go.mod h1:ijPqXp5P6IRRByFVVg9DY8P5HkxkHE5ARIa+86aXPf4=
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.59.0 h1:CV7UdSGJt/Ao6Gp4CXckLxVRRsRgDHoI8XjbL3PDl8s=
go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.59.0/go.mod h1:FRmFuRJfag1IZ2dPkHnEoSFVgTVPUd2qf5Vi69hLb8I=
go.opentelemetry.io/otel v1.34.0 h1:zRLXxLCgL1WyKsPVrgbSdMN4c0FMkDAskSTQP+0hdUY=
go.opentelemetry.io/otel v1.34.0/go.mod h1:OWFPOQ+h4G8xpyjgqo4SxJYdDQ/qmRH+wivy7zzx9oI=
go.opentelemetry.io/otel/metric v1.34.0 h1:+eTR3U0MyfWjRDhmFMxe2SsW64QrZ84AOhvqS7Y+PoQ=
go.opentelemetry.io/otel/metric v1.34.0/go.mod h1:CEDrp0fy2D0MvkXE+dPV7cMi8tWZwX3dmaIhwPOaqHE=
go.opentelemetry.io/otel/sdk v1.34.0 h1:95zS4k/2GOy069d321O8jWgYsW3MzVV+KuSPKp7Wr1A=
go.opentelemetry.io/otel/sdk v1.34.0/go.mod h1:0e/pNiaMAqaykJGKbi+tSjWfNNHMTxoC9qANsCzbyxU=
go.opentelemetry.io/otel/sdk/metric v1.32.0 h1:rZvFnvmvawYb0alrYkjraqJq0Z4ZUJAiyYCU9snn1CU=
go.opentelemetry.io/otel/sdk/metric v1.32.0/go.mod h1:PWeZlq0zt9YkYAp3gjKZ0eicRYvOh1Gd+X99x6GHpCQ=
go.opentelemetry.io/otel/trace v1.34.0 h1:+ouXS2V8Rd4hp4580a8q23bg0azF2nI8cqLYnC8mh/k=
go.opentelemetry.io/otel/trace v1.34.0/go.mod h1:Svm7lSjQD7kG7KJ/MUHPVXSDGz2OX4h0M2jHBhmSfRE=
golang.org/x/arch v0.14.0 h1:z9JUEZWr8x4rR0OU6c4/4t6E6jOZ8/QBS2bBYBm4tx4=
golang.org/x/arch v0.14.0/go.mod h1:FEVrYAQjsQXMVJ1nsMoVVXPZg6p2JE2mx8psSWTDQys=
golang.org/x/crypto v0.33.0 h1:IOBPskki6Lysi0lo9qQvbxiQ+FvsCC/YWOecCHAixus=
golang.org/x/crypto v0.33.0/go.mod h1:bVdXmD7IV/4GdElGPozy6U7lWdRXA4qyRVGJV57uQ5M=
golang.org/x/exp v0.0.0-20250210185358-939b2ce775ac h1:l5+whBCLH3iH2ZNHYLbAe58bo7yrN4mVcnkHDYz5vvs=
golang.org/x/exp v0.0.0-20250210185358-939b2ce775ac/go.mod h1:hH+7mtFmImwwcMvScyxUhjuVHR3HGaDPMn9rMSUUbxo=
golang.org/x/image v0.24.0 h1:AN7zRgVsbvmTfNyqIbbOraYL8mSwcKncEj8ofjgzcMQ=
golang.org/x/image v0.24.0/go.mod h1:4b/ITuLfqYq1hqZcjofwctIhi7sZh2WaCjvsBNjjya8=
golang.org/x/mod v0.23.0 h1:Zb7khfcRGKk+kqfxFaP5tZqCnDZMjC5VtUBs87Hr6QM=
golang.org/x/mod v0.23.0/go.mod h1:6SkKJ3Xj0I0BrPOZoBy3bdMptDDU9oJrpohJ3eWZ1fY=
golang.org/x/net v0.35.0 h1:T5GQRQb2y08kTAByq9L4/bz8cipCdA8FbRTXewonqY8=
golang.org/x/net v0.35.0/go.mod h1:EglIi67kWsHKlRzzVMUD93VMSWGFOMSZgxFjparz1Qk=
golang.org/x/oauth2 v0.26.0 h1:afQXWNNaeC4nvZ0Ed9XvCCzXM6UHJG7iCg0W4fPqSBE=
golang.org/x/oauth2 v0.26.0/go.mod h1:XYTD2NtWslqkgxebSiOHnXEap4TF09sJSc7H1sXbhtI=
golang.org/x/sync v0.11.0 h1:GGz8+XQP4FvTTrjZPzNKTMFtSXH80RAzG+5ghFPgK9w=
golang.org/x/sync v0.11.0/go.mod h1:Czt+wKu1gCyEFDUtn0jG5QVvpJ6rzVqr5aXyt9drQfk=
golang.org/x/sys v0.0.0-20220715151400-c0bba94af5f8/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.28.0 h1:Fksou7UEQUWlKvIdsqzJmUmCX3cZuD2+P3XyyzwMhlA=
golang.org/x/sys v0.28.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ=
golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ=
golang.org/x/text v0.21.0 h1:zyQAAkrwaneQ066sspRyJaG9VNi/YJ1NfzcGB3hZ/qo=
golang.org/x/text v0.21.0/go.mod h1:4IBbMaMmOPCJ8SecivzSH54+73PCFmPWxNTLm+vZkEQ=
golang.org/x/time v0.8.0 h1:9i3RxcPv3PZnitoVGMPDKZSq1xW1gK1Xy3ArNOGZfEg=
golang.org/x/time v0.8.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM=
golang.org/x/tools v0.0.0-20180917221912-90fa682c2a6e/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
golang.org/x/tools v0.0.0-20190114222345-bf090417da8b/go.mod h1:n7NCudcB/nEzxVGmLbDWY5pfWTLqBcC2KZ6jyYvM4mQ=
golang.org/x/tools v0.0.0-20190226205152-f727befe758c/go.mod h1:9Yl7xja0Znq3iFh3HoIrodX9oNMXvdceNzlUR8zjMvY=
golang.org/x/tools v0.0.0-20190311212946-11955173bddd/go.mod h1:LCzVGOaR6xXOjkQ3onu1FJEFr0SW1gC7cKk1uF8kGRs=
golang.org/x/tools v0.0.0-20190524140312-2c0ae7006135/go.mod h1:RgjU9mgBXZiqYHBnxXauZ1Gv1EHHAz9KjViQ78xBX0Q=
golang.org/x/tools v0.28.0 h1:WuB6qZ4RPCQo5aP3WdKZS7i595EdWqWR8vqJTlwTVK8=
golang.org/x/tools v0.28.0/go.mod h1:dcIOrVd3mfQKTgrDVQHqCPMWy6lnhfhtX3hLXYVLfRw=
golang.org/x/xerrors v0.0.0-20191204190536-9bdfabe68543/go.mod h1:I/5z698sn9Ka8TeJc9MKroUUfqBBauWjQqLJ2OPfmY0=
google.golang.org/api v0.210.0 h1:HMNffZ57OoZCRYSbdWVRoqOa8V8NIHLL0CzdBPLztWk=
google.golang.org/api v0.210.0/go.mod h1:B9XDZGnx2NtyjzVkOVTGrFSAVZgPcbedzKg/gTLwqBs=
google.golang.org/appengine v1.1.0/go.mod h1:EbEs0AVv82hx2wNQdGPgUI5lhzA/G0D9YwlJXL52JkM=
google.golang.org/appengine v1.4.0/go.mod h1:xpcJRLb0r/rnEns0DIKYYv+WjYCduHsrkT7/EB5XEv4=
google.golang.org/genproto v0.0.0-20180817151627-c66870c02cf8/go.mod h1:JiN7NxoALGmiZfu7CAH4rXhgtRTLTxftemlI0sWmxmc=
google.golang.org/genproto v0.0.0-20190819201941-24fa4b261c55/go.mod h1:DMBHOl98Agz4BDEuKkezgsaosCRResVns1a3J2ZsMNc=
google.golang.org/genproto v0.0.0-20200526211855-cb27e3aa2013/go.mod h1:NbSheEEYHJ7i3ixzK3sjbqSGDJWnxyFXZblF3eUsNvo=
google.golang.org/genproto/googleapis/api v0.0.0-20241209162323-e6fa225c2576 h1:CkkIfIt50+lT6NHAVoRYEyAvQGFM7xEwXUUywFvEb3Q=
google.golang.org/genproto/googleapis/api v0.0.0-20241209162323-e6fa225c2576/go.mod h1:1R3kvZ1dtP3+4p4d3G8uJ8rFk/fWlScl38vanWACI08=
google.golang.org/genproto/googleapis/rpc v0.0.0-20241209162323-e6fa225c2576 h1:8ZmaLZE4XWrtU3MyClkYqqtl6Oegr3235h7jxsDyqCY=
google.golang.org/genproto/googleapis/rpc v0.0.0-20241209162323-e6fa225c2576/go.mod h1:5uTbfoYQed2U9p3KIj2/Zzm02PYhndfdmML0qC3q3FU=
google.golang.org/grpc v1.19.0/go.mod h1:mqu4LbDTu4XGKhr4mRzUsmM4RtVoemTSY81AxZiDr8c=
google.golang.org/grpc v1.23.0/go.mod h1:Y5yQAOtifL1yxbo5wqy6BxZv8vAUGQwXBOALyacEbxg=
google.golang.org/grpc v1.25.1/go.mod h1:c3i+UQWmh7LiEpx4sFZnkU36qjEYZ0imhYfXVyQciAY=
google.golang.org/grpc v1.27.0/go.mod h1:qbnxyOmOxrQa7FizSgH+ReBfzJrCY1pSN7KXBS8abTk=
google.golang.org/grpc v1.33.2/go.mod h1:JMHMWHQWaTccqQQlmk3MJZS+GWXOdAesneDmEnv2fbc=
google.golang.org/grpc v1.68.1 h1:oI5oTa11+ng8r8XMMN7jAOmWfPZWbYpCFaMUTACxkM0=
google.golang.org/grpc v1.68.1/go.mod h1:+q1XYFJjShcqn0QZHvCyeR4CXPA+llXIeUIfIe00waw=
google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8=
google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0=
google.golang.org/protobuf v0.0.0-20200228230310-ab0ca4ff8a60/go.mod h1:cfTl7dwQJ+fmap5saPgwCLgHXTUD7jkjRqWcaiX5VyM=
google.golang.org/protobuf v1.20.1-0.20200309200217-e05f789c0967/go.mod h1:A+miEFZTKqfCUM6K7xSMQL9OKL/b6hQv+e19PK+JZNE=
google.golang.org/protobuf v1.21.0/go.mod h1:47Nbq4nVaFHyn7ilMalzfO3qCViNmqZ2kzikPIcrTAo=
google.golang.org/protobuf v1.22.0/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2l/sGQquU=
google.golang.org/protobuf v1.23.0/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2l/sGQquU=
google.golang.org/protobuf v1.23.1-0.20200526195155-81db48ad09cc/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2l/sGQquU=
google.golang.org/protobuf v1.25.0/go.mod h1:9JNX74DMeImyA3h4bdi1ymwjUzf21/xIlbajtzgsN7c=
google.golang.org/protobuf v1.35.2 h1:8Ar7bF+apOIoThw1EdZl0p1oWvMqTHmpA2fRTyZO8io=
google.golang.org/protobuf v1.35.2/go.mod h1:9fA7Ob0pmnwhb644+1+CVWFRbNajQ6iRojtC/QF5bRE=
golang.org/x/sys v0.30.0 h1:QjkSwP/36a20jFYWkSue1YwXzLmsV5Gfq7Eiy72C1uc=
golang.org/x/sys v0.30.0/go.mod h1:/VUhepiaJMQUp4+oa/7Zr1D23ma6VTLIYjOOTFZPUcA=
golang.org/x/text v0.22.0 h1:bofq7m3/HAFvbF51jz3Q9wLg3jkvSPuiZu/pD1XwgtM=
golang.org/x/text v0.22.0/go.mod h1:YRoo4H8PVmsu+E3Ou7cqLVH8oXWIHVoX0jqUWALQhfY=
golang.org/x/time v0.10.0 h1:3usCWA8tQn0L8+hFJQNgzpWbd89begxN66o1Ojdn5L4=
golang.org/x/time v0.10.0/go.mod h1:3BpzKBy/shNhVucY/MWOyx10tF3SFh9QdLuxbVysPQM=
golang.org/x/tools v0.30.0 h1:BgcpHewrV5AUp2G9MebG4XPFI1E2W41zU1SaqVA9vJY=
golang.org/x/tools v0.30.0/go.mod h1:c347cR/OJfw5TI+GfX7RUPNMdDRRbjvYTS0jPyvsVtY=
google.golang.org/api v0.220.0 h1:3oMI4gdBgB72WFVwE1nerDD8W3HUOS4kypK6rRLbGns=
google.golang.org/api v0.220.0/go.mod h1:26ZAlY6aN/8WgpCzjPNy18QpYaz7Zgg1h0qe1GkZEmY=
google.golang.org/genproto/googleapis/api v0.0.0-20250207221924-e9438ea467c6 h1:L9JNMl/plZH9wmzQUHleO/ZZDSN+9Gh41wPczNy+5Fk=
google.golang.org/genproto/googleapis/api v0.0.0-20250207221924-e9438ea467c6/go.mod h1:iYONQfRdizDB8JJBybql13nArx91jcUk7zCXEsOofM4=
google.golang.org/genproto/googleapis/rpc v0.0.0-20250207221924-e9438ea467c6 h1:2duwAxN2+k0xLNpjnHTXoMUgnv6VPSp5fiqTuwSxjmI=
google.golang.org/genproto/googleapis/rpc v0.0.0-20250207221924-e9438ea467c6/go.mod h1:8BS3B93F/U1juMFq9+EDk+qOT5CO1R9IzXxG3PTqiRk=
google.golang.org/grpc v1.70.0 h1:pWFv03aZoHzlRKHWicjsZytKAiYCtNS0dHbXnIdq7jQ=
google.golang.org/grpc v1.70.0/go.mod h1:ofIJqVKDXx/JiXrwr2IG4/zwdH9txy3IlF40RmcJSQw=
google.golang.org/protobuf v1.36.5 h1:tPhr+woSbjfYvY6/GPufUoYizxw1cF/yFoxJ2fmpwlM=
google.golang.org/protobuf v1.36.5/go.mod h1:9fA7Ob0pmnwhb644+1+CVWFRbNajQ6iRojtC/QF5bRE=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk=
gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q=
@@ -320,30 +250,28 @@ gorm.io/driver/postgres v1.5.11/go.mod h1:DX3GReXH+3FPWGrrgffdvCk3DQ1dwDPdmbenSk
gorm.io/gorm v1.25.7/go.mod h1:hbnx/Oo0ChWMn1BIhpy1oYozzpM15i4YPuHDmfYtwg8=
gorm.io/gorm v1.25.12 h1:I0u8i2hWQItBq1WfE0o2+WuL9+8L21K9e2HHSTE/0f8=
gorm.io/gorm v1.25.12/go.mod h1:xh7N7RHfYlNc5EmcI/El95gXusucDrQnHXe0+CgWcLQ=
honnef.co/go/tools v0.0.0-20190102054323-c2f93a96b099/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4=
honnef.co/go/tools v0.0.0-20190523083050-ea95bdfd59fc/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4=
modernc.org/cc/v4 v4.23.1 h1:WqJoPL3x4cUufQVHkXpXX7ThFJ1C4ik80i2eXEXbhD8=
modernc.org/cc/v4 v4.23.1/go.mod h1:HM7VJTZbUCR3rV8EYBi9wxnJ0ZBRiGE5OeGXNA0IsLQ=
modernc.org/ccgo/v4 v4.23.1 h1:N49a7JiWGWV7lkPE4yYcvjkBGZQi93/JabRYjdWmJXc=
modernc.org/ccgo/v4 v4.23.1/go.mod h1:JoIUegEIfutvoWV/BBfDFpPpfR2nc3U0jKucGcbmwDU=
modernc.org/cc/v4 v4.24.4 h1:TFkx1s6dCkQpd6dKurBNmpo+G8Zl4Sq/ztJ+2+DEsh0=
modernc.org/cc/v4 v4.24.4/go.mod h1:uVtb5OGqUKpoLWhqwNQo/8LwvoiEBLvZXIQ/SmO6mL0=
modernc.org/ccgo/v4 v4.23.16 h1:Z2N+kk38b7SfySC1ZkpGLN2vthNJP1+ZzGZIlH7uBxo=
modernc.org/ccgo/v4 v4.23.16/go.mod h1:nNma8goMTY7aQZQNTyN9AIoJfxav4nvTnvKThAeMDdo=
modernc.org/fileutil v1.3.0 h1:gQ5SIzK3H9kdfai/5x41oQiKValumqNTDXMvKo62HvE=
modernc.org/fileutil v1.3.0/go.mod h1:XatxS8fZi3pS8/hKG2GH/ArUogfxjpEKs3Ku3aK4JyQ=
modernc.org/gc/v2 v2.5.0 h1:bJ9ChznK1L1mUtAQtxi0wi5AtAs5jQuw4PrPHO5pb6M=
modernc.org/gc/v2 v2.5.0/go.mod h1:wzN5dK1AzVGoH6XOzc3YZ+ey/jPgYHLuVckd62P0GYU=
modernc.org/libc v1.61.4 h1:wVyqEx6tlltte9lPTjq0kDAdtdM9c4JH8rU6M1ZVawA=
modernc.org/libc v1.61.4/go.mod h1:VfXVuM/Shh5XsMNrh3C6OkfL78G3loa4ZC/Ljv9k7xc=
modernc.org/mathutil v1.6.0 h1:fRe9+AmYlaej+64JsEEhoWuAYBkOtQiMEU7n/XgfYi4=
modernc.org/mathutil v1.6.0/go.mod h1:Ui5Q9q1TR2gFm0AQRqQUaBWFLAhQpCwNcuhBOSedWPo=
modernc.org/memory v1.8.0 h1:IqGTL6eFMaDZZhEWwcREgeMXYwmW83LYW8cROZYkg+E=
modernc.org/memory v1.8.0/go.mod h1:XPZ936zp5OMKGWPqbD3JShgd/ZoQ7899TUuQqxY+peU=
modernc.org/opt v0.1.3 h1:3XOZf2yznlhC+ibLltsDGzABUGVx8J6pnFMS3E4dcq4=
modernc.org/opt v0.1.3/go.mod h1:WdSiB5evDcignE70guQKxYUl14mgWtbClRi5wmkkTX0=
modernc.org/sortutil v1.2.0 h1:jQiD3PfS2REGJNzNCMMaLSp/wdMNieTbKX920Cqdgqc=
modernc.org/sortutil v1.2.0/go.mod h1:TKU2s7kJMf1AE84OoiGppNHJwvB753OYfNl2WRb++Ss=
modernc.org/sqlite v1.34.2 h1:J9n76TPsfYYkFkZ9Uy1QphILYifiVEwwOT7yP5b++2Y=
modernc.org/sqlite v1.34.2/go.mod h1:dnR723UrTtjKpoHCAMN0Q/gZ9MT4r+iRvIBb9umWFkU=
modernc.org/strutil v1.2.0 h1:agBi9dp1I+eOnxXeiZawM8F4LawKv4NzGWSaLfyeNZA=
modernc.org/strutil v1.2.0/go.mod h1:/mdcBmfOibveCTBxUl5B5l6W+TTH1FXPLHZE6bTosX0=
modernc.org/gc/v2 v2.6.3 h1:aJVhcqAte49LF+mGveZ5KPlsp4tdGdAOT4sipJXADjw=
modernc.org/gc/v2 v2.6.3/go.mod h1:YgIahr1ypgfe7chRuJi2gD7DBQiKSLMPgBQe9oIiito=
modernc.org/libc v1.61.12 h1:Fsnh0A7XLXylYNwIOJmKux9PhnfrIvMaMnjuyJ1t/f4=
modernc.org/libc v1.61.12/go.mod h1:8F/uJWL/3nNil0Lgt1Dpz+GgkApWh04N3el3hxJcA6E=
modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU=
modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg=
modernc.org/memory v1.8.2 h1:cL9L4bcoAObu4NkxOlKWBWtNHIsnnACGF/TbqQ6sbcI=
modernc.org/memory v1.8.2/go.mod h1:ZbjSvMO5NQ1A2i3bWeDiVMxIorXwdClKE/0SZ+BMotU=
modernc.org/opt v0.1.4 h1:2kNGMRiUjrp4LcaPuLY2PzUfqM/w9N23quVwhKt5Qm8=
modernc.org/opt v0.1.4/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns=
modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w=
modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE=
modernc.org/sqlite v1.34.5 h1:Bb6SR13/fjp15jt70CL4f18JIN7p7dnMExd+UFnF15g=
modernc.org/sqlite v1.34.5/go.mod h1:YLuNmX9NKs8wRNK2ko1LW1NGYcc9FkBO69JOt1AR9JE=
modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0=
modernc.org/strutil v1.2.1/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A=
modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y=
modernc.org/token v1.1.0/go.mod h1:UGzOrNV1mAFSEB63lOFHIpNRUVMvYTc6yu1SMY/XTDM=
nullprogram.com/x/optparse v1.0.0/go.mod h1:KdyPE+Igbe0jQUrVfMqDMeJQIJZEuyV7pjYmp6pbG50=
+32 -24
View File
@@ -17,11 +17,11 @@ import (
_ "github.com/joho/godotenv/autoload"
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/common/balance"
"github.com/labring/sealos/service/aiproxy/common/client"
"github.com/labring/sealos/service/aiproxy/common/config"
"github.com/labring/sealos/service/aiproxy/common/consume"
"github.com/labring/sealos/service/aiproxy/controller"
"github.com/labring/sealos/service/aiproxy/middleware"
"github.com/labring/sealos/service/aiproxy/model"
relaycontroller "github.com/labring/sealos/service/aiproxy/relay/controller"
"github.com/labring/sealos/service/aiproxy/router"
log "github.com/sirupsen/logrus"
)
@@ -39,12 +39,7 @@ func initializeServices() error {
return err
}
if err := initializeCaches(); err != nil {
return err
}
client.Init()
return nil
return initializeCaches()
}
func initializeBalance() error {
@@ -76,7 +71,7 @@ func setLog(l *log.Logger) {
l.SetOutput(os.Stdout)
stdlog.SetOutput(l.Writer())
log.SetFormatter(&log.TextFormatter{
l.SetFormatter(&log.TextFormatter{
ForceColors: true,
DisableColors: false,
ForceQuote: config.DebugEnabled,
@@ -105,20 +100,16 @@ func initializeDatabases() error {
}
func initializeCaches() error {
if err := model.InitOptionMap(); err != nil {
if err := model.InitOption2DB(); err != nil {
return err
}
if err := model.InitModelConfigCache(); err != nil {
return err
}
return model.InitChannelCache()
return model.InitModelConfigAndChannelCache()
}
func startSyncServices(ctx context.Context, wg *sync.WaitGroup) {
wg.Add(3)
wg.Add(2)
go model.SyncOptions(ctx, wg, time.Second*5)
go model.SyncChannelCache(ctx, wg, time.Second*5)
go model.SyncModelConfigCache(ctx, wg, time.Second*5)
go model.SyncModelConfigAndChannelCache(ctx, wg, time.Second*10)
}
func setupHTTPServer() (*http.Server, *gin.Engine) {
@@ -126,9 +117,9 @@ func setupHTTPServer() (*http.Server, *gin.Engine) {
w := log.StandardLogger().Writer()
server.
Use(middleware.NewLog(log.StandardLogger())).
Use(gin.RecoveryWithWriter(w)).
Use(middleware.RequestID)
Use(middleware.NewLog(log.StandardLogger())).
Use(middleware.RequestID, middleware.CORS())
router.SetRouter(server)
port := os.Getenv("PORT")
@@ -143,6 +134,16 @@ func setupHTTPServer() (*http.Server, *gin.Engine) {
}, server
}
func autoTestBannedModels() {
log.Info("auto test banned models start")
ticker := time.NewTicker(time.Second * 15)
defer ticker.Stop()
for range ticker.C {
controller.AutoTestBannedModels()
}
}
func main() {
if err := initializeServices(); err != nil {
log.Fatal("failed to initialize services: " + err.Error())
@@ -169,18 +170,25 @@ func main() {
}
}()
<-ctx.Done()
log.Info("shutting down server...")
log.Info("max wait time: 120s")
go autoTestBannedModels()
shutdownCtx, cancel := context.WithTimeout(context.Background(), 120*time.Second)
<-ctx.Done()
shutdownCtx, cancel := context.WithTimeout(context.Background(), 600*time.Second)
defer cancel()
log.Info("shutting down http server...")
log.Info("max wait time: 600s")
if err := srv.Shutdown(shutdownCtx); err != nil {
log.Error("server forced to shutdown: " + err.Error())
} else {
log.Info("server shutdown successfully")
}
relaycontroller.ConsumeWaitGroup.Wait()
log.Info("shutting down consumer...")
consume.Wait()
log.Info("shutting down sync services...")
wg.Wait()
log.Info("server exiting")
+101 -33
View File
@@ -4,7 +4,6 @@ import (
"fmt"
"net/http"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/common/config"
@@ -42,12 +41,18 @@ func AdminAuth(c *gin.Context) {
c.Abort()
return
}
group := c.Param("group")
if group != "" {
log := GetLogger(c)
log.Data["gid"] = group
}
c.Next()
}
func TokenAuth(c *gin.Context) {
log := GetLogger(c)
ctx := c.Request.Context()
key := c.Request.Header.Get("Authorization")
key = strings.TrimPrefix(
strings.TrimPrefix(key, "Bearer "),
@@ -55,18 +60,29 @@ func TokenAuth(c *gin.Context) {
)
parts := strings.Split(key, "-")
key = parts[0]
token, err := model.ValidateAndGetToken(key)
if err != nil {
abortWithMessage(c, http.StatusUnauthorized, err.Error())
return
var token *model.TokenCache
var useInternalToken bool
if config.GetInternalToken() != "" && config.GetInternalToken() == key || config.AdminKey != "" && config.AdminKey == key {
token = &model.TokenCache{}
useInternalToken = true
} else {
var err error
token, err = model.ValidateAndGetToken(key)
if err != nil {
abortLogWithMessage(c, http.StatusUnauthorized, err.Error())
return
}
}
SetLogTokenFields(log.Data, token)
SetLogTokenFields(log.Data, token, useInternalToken)
if token.Subnet != "" {
if ok, err := network.IsIPInSubnets(c.ClientIP(), token.Subnet); err != nil {
abortWithMessage(c, http.StatusInternalServerError, err.Error())
abortLogWithMessage(c, http.StatusInternalServerError, err.Error())
return
} else if !ok {
abortWithMessage(c, http.StatusForbidden,
abortLogWithMessage(c, http.StatusForbidden,
fmt.Sprintf("token (%s[%d]) can only be used in the specified subnet: %s, current ip: %s",
token.Name,
token.ID,
@@ -77,47 +93,86 @@ func TokenAuth(c *gin.Context) {
return
}
}
group, err := model.CacheGetGroup(token.Group)
if err != nil {
abortWithMessage(c, http.StatusInternalServerError, err.Error())
return
}
SetLogGroupFields(log.Data, group)
if len(token.Models) == 0 {
token.Models = model.CacheGetEnabledModels()
}
if group.QPM <= 0 {
group.QPM = config.GetDefaultGroupQPM()
}
if group.QPM > 0 {
ok := ForceRateLimit(ctx, "group_qpm:"+group.ID, int(group.QPM), time.Minute)
if !ok {
abortWithMessage(c, http.StatusTooManyRequests,
group.ID+" is requesting too frequently",
)
var group *model.GroupCache
if useInternalToken {
group = &model.GroupCache{
Status: model.GroupStatusInternal,
}
} else {
var err error
group, err = model.CacheGetGroup(token.Group)
if err != nil {
abortLogWithMessage(c, http.StatusInternalServerError, fmt.Sprintf("failed to get group: %v", err))
return
}
if group.Status != model.GroupStatusEnabled && group.Status != model.GroupStatusInternal {
abortLogWithMessage(c, http.StatusForbidden, "group is disabled")
return
}
}
SetLogGroupFields(log.Data, group)
modelCaches := model.LoadModelCaches()
storeTokenModels(token, modelCaches)
c.Set(ctxkey.Group, group)
c.Set(ctxkey.Token, token)
c.Set(ctxkey.ModelCaches, modelCaches)
c.Next()
}
func GetGroup(c *gin.Context) *model.GroupCache {
return c.MustGet(ctxkey.Group).(*model.GroupCache)
}
func GetToken(c *gin.Context) *model.TokenCache {
return c.MustGet(ctxkey.Token).(*model.TokenCache)
}
func GetModelCaches(c *gin.Context) *model.ModelCaches {
return c.MustGet(ctxkey.ModelCaches).(*model.ModelCaches)
}
func sliceFilter[T any](s []T, fn func(T) bool) []T {
i := 0
for _, v := range s {
if fn(v) {
s[i] = v
i++
}
}
return s[:i]
}
func storeTokenModels(token *model.TokenCache, modelCaches *model.ModelCaches) {
if len(token.Models) == 0 {
token.Models = modelCaches.EnabledModels
} else {
enabledModelsMap := modelCaches.EnabledModelsMap
token.Models = sliceFilter(token.Models, func(m string) bool {
_, ok := enabledModelsMap[m]
return ok
})
}
}
func SetLogFieldsFromMeta(m *meta.Meta, fields logrus.Fields) {
SetLogRequestIDField(fields, m.RequestID)
SetLogModeField(fields, m.Mode)
SetLogModelFields(fields, m.OriginModelName)
SetLogActualModelFields(fields, m.ActualModelName)
SetLogModelFields(fields, m.OriginModel)
SetLogActualModelFields(fields, m.ActualModel)
if m.IsChannelTest {
SetLogIsChannelTestField(fields, true)
}
SetLogGroupFields(fields, m.Group)
SetLogTokenFields(fields, m.Token)
SetLogTokenFields(fields, m.Token, false)
SetLogChannelFields(fields, m.Channel)
}
@@ -150,17 +205,30 @@ func SetLogRequestIDField(fields logrus.Fields, requestID string) {
}
func SetLogGroupFields(fields logrus.Fields, group *model.GroupCache) {
if group != nil {
if group == nil {
return
}
if group.ID != "" {
fields["gid"] = group.ID
}
}
func SetLogTokenFields(fields logrus.Fields, token *model.TokenCache) {
if token != nil {
func SetLogTokenFields(fields logrus.Fields, token *model.TokenCache, internal bool) {
if token == nil {
return
}
if token.ID > 0 {
fields["tid"] = token.ID
}
if token.Name != "" {
fields["tname"] = token.Name
}
if token.Key != "" {
fields["key"] = maskTokenKey(token.Key)
}
if internal {
fields["internal"] = "true"
}
}
func maskTokenKey(key string) string {
+226 -27
View File
@@ -1,47 +1,207 @@
package middleware
import (
"context"
"errors"
"fmt"
"net/http"
"slices"
"strconv"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/common/balance"
"github.com/labring/sealos/service/aiproxy/common/config"
"github.com/labring/sealos/service/aiproxy/common/consume"
"github.com/labring/sealos/service/aiproxy/common/ctxkey"
"github.com/labring/sealos/service/aiproxy/common/helper"
"github.com/labring/sealos/service/aiproxy/common/rpmlimit"
"github.com/labring/sealos/service/aiproxy/model"
"github.com/labring/sealos/service/aiproxy/relay/meta"
"github.com/labring/sealos/service/aiproxy/relay/relaymode"
log "github.com/sirupsen/logrus"
)
type ModelRequest struct {
Model string `form:"model" json:"model"`
func calculateGroupConsumeLevelRatio(usedAmount float64) float64 {
v := config.GetGroupConsumeLevelRatio()
if len(v) == 0 {
return 1
}
var maxConsumeLevel float64 = -1
var groupConsumeLevelRatio float64
for consumeLevel, ratio := range v {
if usedAmount < consumeLevel {
continue
}
if consumeLevel > maxConsumeLevel {
maxConsumeLevel = consumeLevel
groupConsumeLevelRatio = ratio
}
}
if groupConsumeLevelRatio <= 0 {
groupConsumeLevelRatio = 1
}
return groupConsumeLevelRatio
}
func Distribute(c *gin.Context) {
func getGroupPMRatio(group *model.GroupCache) (float64, float64) {
groupRPMRatio := group.RPMRatio
if groupRPMRatio <= 0 {
groupRPMRatio = 1
}
groupTPMRatio := group.TPMRatio
if groupTPMRatio <= 0 {
groupTPMRatio = 1
}
return groupRPMRatio, groupTPMRatio
}
func GetGroupAdjustedModelConfig(group *model.GroupCache, mc *model.ModelConfig) *model.ModelConfig {
rpm := mc.RPM
tpm := mc.TPM
if group.RPM != nil && group.RPM[mc.Model] > 0 {
rpm = group.RPM[mc.Model]
}
if group.TPM != nil && group.TPM[mc.Model] > 0 {
tpm = group.TPM[mc.Model]
}
rpmRatio, tpmRatio := getGroupPMRatio(group)
groupConsumeLevelRatio := calculateGroupConsumeLevelRatio(group.UsedAmount)
rpm = int64(float64(rpm) * rpmRatio * groupConsumeLevelRatio)
tpm = int64(float64(tpm) * tpmRatio * groupConsumeLevelRatio)
if rpm != mc.RPM || tpm != mc.TPM {
newMc := *mc
newMc.RPM = rpm
newMc.TPM = tpm
return &newMc
}
return mc
}
var (
ErrRequestRateLimitExceeded = errors.New("request rate limit exceeded, please try again later")
ErrRequestTpmLimitExceeded = errors.New("request tpm limit exceeded, please try again later")
)
func checkGroupModelRPMAndTPM(c *gin.Context, group *model.GroupCache, mc *model.ModelConfig) error {
adjustedModelConfig := GetGroupAdjustedModelConfig(group, mc)
if adjustedModelConfig.RPM > 0 {
ok := rpmlimit.ForceRateLimit(
c.Request.Context(),
group.ID,
mc.Model,
adjustedModelConfig.RPM,
time.Minute,
)
if !ok {
return ErrRequestRateLimitExceeded
}
} else if common.RedisEnabled {
_, err := rpmlimit.PushRequest(c.Request.Context(), group.ID, mc.Model, time.Minute)
if err != nil {
log.Errorf("push request error: %s", err.Error())
}
}
if adjustedModelConfig.TPM > 0 {
tpm, err := model.CacheGetGroupModelTPM(group.ID, mc.Model)
if err != nil {
log.Errorf("get group model tpm (%s:%s) error: %s", group.ID, mc.Model, err.Error())
// ignore error
return nil
}
if tpm >= adjustedModelConfig.TPM {
return ErrRequestTpmLimitExceeded
}
}
return nil
}
type GroupBalanceConsumer struct {
GroupBalance float64
Consumer balance.PostGroupConsumer
}
func checkGroupBalance(c *gin.Context, group *model.GroupCache) bool {
var groupBalance float64
var consumer balance.PostGroupConsumer
if group.Status == model.GroupStatusInternal {
groupBalance, consumer, _ = balance.MockGetGroupRemainBalance(c.Request.Context(), *group)
} else {
log := GetLogger(c)
var err error
groupBalance, consumer, err = balance.Default.GetGroupRemainBalance(c.Request.Context(), *group)
if err != nil {
if errors.Is(err, balance.ErrNoRealNameUsedAmountLimit) {
abortLogWithMessage(c, http.StatusForbidden, balance.ErrNoRealNameUsedAmountLimit.Error())
return false
}
log.Errorf("get group (%s) balance error: %v", group.ID, err)
abortWithMessage(c, http.StatusInternalServerError, fmt.Sprintf("get group (%s) balance error", group.ID))
return false
}
log.Data["balance"] = strconv.FormatFloat(groupBalance, 'f', -1, 64)
}
if groupBalance <= 0 {
abortLogWithMessage(c, http.StatusForbidden, fmt.Sprintf("group (%s) balance not enough", group.ID))
return false
}
c.Set(ctxkey.GroupBalance, &GroupBalanceConsumer{
GroupBalance: groupBalance,
Consumer: consumer,
})
return true
}
func NewDistribute(mode int) gin.HandlerFunc {
return func(c *gin.Context) {
distribute(c, mode)
}
}
func distribute(c *gin.Context, mode int) {
if config.GetDisableServe() {
abortWithMessage(c, http.StatusServiceUnavailable, "service is under maintenance")
abortLogWithMessage(c, http.StatusServiceUnavailable, "service is under maintenance")
return
}
log := GetLogger(c)
group := GetGroup(c)
if !checkGroupBalance(c, group) {
return
}
requestModel, err := getRequestModel(c)
if err != nil {
abortWithMessage(c, http.StatusBadRequest, err.Error())
abortLogWithMessage(c, http.StatusBadRequest, err.Error())
return
}
if requestModel == "" {
abortWithMessage(c, http.StatusBadRequest, "no model provided")
abortLogWithMessage(c, http.StatusBadRequest, "no model provided")
return
}
c.Set(ctxkey.OriginalModel, requestModel)
SetLogModelFields(log.Data, requestModel)
token := c.MustGet(ctxkey.Token).(*model.TokenCache)
mc, ok := GetModelCaches(c).ModelConfig.GetModelConfig(requestModel)
if !ok {
abortLogWithMessage(c, http.StatusServiceUnavailable, requestModel+" is not available")
return
}
c.Set(ctxkey.ModelConfig, mc)
token := GetToken(c)
if len(token.Models) == 0 || !slices.Contains(token.Models, requestModel) {
abortWithMessage(c,
abortLogWithMessage(c,
http.StatusForbidden,
fmt.Sprintf("token (%s[%d]) has no permission to use model: %s",
token.Name, token.ID, requestModel,
@@ -49,32 +209,71 @@ func Distribute(c *gin.Context) {
)
return
}
channel, err := model.CacheGetRandomSatisfiedChannel(requestModel)
if err != nil {
abortWithMessage(c, http.StatusServiceUnavailable, requestModel+" is not available")
if err := checkGroupModelRPMAndTPM(c, group, mc); err != nil {
errMsg := err.Error()
consume.AsyncConsume(
nil,
http.StatusTooManyRequests,
nil,
NewMetaByContext(c, nil, mc.Model, mode),
0,
0,
errMsg,
c.ClientIP(),
nil,
)
abortLogWithMessage(c, http.StatusTooManyRequests, errMsg)
return
}
c.Set(string(ctxkey.OriginalModel), requestModel)
ctx := context.WithValue(c.Request.Context(), ctxkey.OriginalModel, requestModel)
c.Request = c.Request.WithContext(ctx)
c.Set(ctxkey.Channel, channel)
c.Next()
}
func NewMetaByContext(c *gin.Context) *meta.Meta {
channel := c.MustGet(ctxkey.Channel).(*model.Channel)
originalModel := c.MustGet(string(ctxkey.OriginalModel)).(string)
requestID := c.GetString(string(helper.RequestIDKey))
group := c.MustGet(ctxkey.Group).(*model.GroupCache)
token := c.MustGet(ctxkey.Token).(*model.TokenCache)
func GetOriginalModel(c *gin.Context) string {
return c.GetString(ctxkey.OriginalModel)
}
func GetModelConfig(c *gin.Context) *model.ModelConfig {
return c.MustGet(ctxkey.ModelConfig).(*model.ModelConfig)
}
func NewMetaByContext(c *gin.Context, channel *model.Channel, modelName string, mode int) *meta.Meta {
requestID := GetRequestID(c)
group := GetGroup(c)
token := GetToken(c)
return meta.NewMeta(
channel,
relaymode.GetByPath(c.Request.URL.Path),
originalModel,
mode,
modelName,
GetModelConfig(c),
meta.WithRequestID(requestID),
meta.WithGroup(group),
meta.WithToken(token),
meta.WithEndpoint(c.Request.URL.Path),
)
}
type ModelRequest struct {
Model string `form:"model" json:"model"`
}
func getRequestModel(c *gin.Context) (string, error) {
path := c.Request.URL.Path
switch {
case strings.HasPrefix(path, "/v1/audio/transcriptions"),
strings.HasPrefix(path, "/v1/audio/translations"):
return c.Request.FormValue("model"), nil
case strings.HasPrefix(path, "/v1/engines") && strings.HasSuffix(path, "/embeddings"):
// /engines/:model/embeddings
return c.Param("model"), nil
default:
var modelRequest ModelRequest
err := common.UnmarshalBodyReusable(c.Request, &modelRequest)
if err != nil {
return "", fmt.Errorf("get request model failed: %w", err)
}
return modelRequest.Model, nil
}
}
-101
View File
@@ -1,101 +0,0 @@
package middleware
import (
"context"
"net/http"
"time"
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/common/config"
log "github.com/sirupsen/logrus"
)
var inMemoryRateLimiter common.InMemoryRateLimiter
// 1. 使用Redis列表存储请求时间戳
// 2. 列表长度代表当前窗口内的请求数
// 3. 如果请求数未达到限制,直接添加新请求并返回成功
// 4. 如果达到限制,则检查最老的请求是否已经过期
// 5. 如果最老的请求已过期,移除它并添加新请求,否则拒绝新请求
// 6. 通过EXPIRE命令设置键的过期时间,自动清理过期数据
var luaScript = `
local key = KEYS[1]
local max_requests = tonumber(ARGV[1])
local window = tonumber(ARGV[2])
local current_time = tonumber(ARGV[3])
local count = redis.call('LLEN', key)
if count < max_requests then
redis.call('LPUSH', key, current_time)
redis.call('PEXPIRE', key, window)
return 1
else
local oldest = redis.call('LINDEX', key, -1)
if current_time - tonumber(oldest) >= window then
redis.call('LPUSH', key, current_time)
redis.call('LTRIM', key, 0, max_requests - 1)
redis.call('PEXPIRE', key, window)
return 1
else
return 0
end
end
`
func redisRateLimitRequest(ctx context.Context, key string, maxRequestNum int, duration time.Duration) (bool, error) {
rdb := common.RDB
currentTime := time.Now().UnixMilli()
result, err := rdb.Eval(ctx, luaScript, []string{key}, maxRequestNum, duration.Milliseconds(), currentTime).Int64()
if err != nil {
return false, err
}
return result == 1, nil
}
func RateLimit(ctx context.Context, key string, maxRequestNum int, duration time.Duration) (bool, error) {
if maxRequestNum == 0 {
return true, nil
}
if common.RedisEnabled {
return redisRateLimitRequest(ctx, key, maxRequestNum, duration)
}
return MemoryRateLimit(ctx, key, maxRequestNum, duration), nil
}
// ignore redis error
func ForceRateLimit(ctx context.Context, key string, maxRequestNum int, duration time.Duration) bool {
if maxRequestNum == 0 {
return true
}
if common.RedisEnabled {
ok, err := redisRateLimitRequest(ctx, key, maxRequestNum, duration)
if err == nil {
return ok
}
log.Error("rate limit error: " + err.Error())
}
return MemoryRateLimit(ctx, key, maxRequestNum, duration)
}
func MemoryRateLimit(_ context.Context, key string, maxRequestNum int, duration time.Duration) bool {
// It's safe to call multi times.
inMemoryRateLimiter.Init(config.RateLimitKeyExpirationDuration)
return inMemoryRateLimiter.Request(key, maxRequestNum, duration)
}
func GlobalAPIRateLimit(c *gin.Context) {
globalAPIRateLimitNum := config.GetGlobalAPIRateLimitNum()
if globalAPIRateLimitNum <= 0 {
c.Next()
return
}
ok := ForceRateLimit(c.Request.Context(), "global_qpm", int(globalAPIRateLimitNum), time.Minute)
if !ok {
c.Status(http.StatusTooManyRequests)
c.Abort()
return
}
c.Next()
}
+21 -6
View File
@@ -1,15 +1,30 @@
package middleware
import (
"strconv"
"time"
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/common/helper"
"github.com/labring/sealos/service/aiproxy/common/ctxkey"
"github.com/labring/sealos/service/aiproxy/common/random"
)
func RequestID(c *gin.Context) {
id := helper.GenRequestID()
c.Set(string(helper.RequestIDKey), id)
c.Header(string(helper.RequestIDKey), id)
func GenRequestID() string {
return strconv.FormatInt(time.Now().UnixMilli(), 10) + random.GetRandomNumberString(4)
}
func SetRequestID(c *gin.Context, id string) {
c.Set(ctxkey.RequestID, id)
c.Header(ctxkey.RequestID, id)
log := GetLogger(c)
SetLogRequestIDField(log.Data, id)
c.Next()
}
func GetRequestID(c *gin.Context) string {
return c.GetString(ctxkey.RequestID)
}
func RequestID(c *gin.Context) {
id := GenRequestID()
SetRequestID(c, id)
}
+10 -23
View File
@@ -2,11 +2,8 @@ package middleware
import (
"fmt"
"strings"
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/common/helper"
"github.com/labring/sealos/service/aiproxy/relay/model"
)
@@ -14,31 +11,21 @@ const (
ErrorTypeAIPROXY = "aiproxy_error"
)
func abortWithMessage(c *gin.Context, statusCode int, message string) {
func MessageWithRequestID(message string, id string) string {
return fmt.Sprintf("%s (aiproxy: %s)", message, id)
}
func abortLogWithMessage(c *gin.Context, statusCode int, message string) {
GetLogger(c).Error(message)
abortWithMessage(c, statusCode, message)
}
func abortWithMessage(c *gin.Context, statusCode int, message string) {
c.JSON(statusCode, gin.H{
"error": &model.Error{
Message: helper.MessageWithRequestID(message, c.GetString(string(helper.RequestIDKey))),
Message: MessageWithRequestID(message, GetRequestID(c)),
Type: ErrorTypeAIPROXY,
},
})
c.Abort()
}
func getRequestModel(c *gin.Context) (string, error) {
path := c.Request.URL.Path
switch {
case strings.HasPrefix(path, "/v1/audio/transcriptions"), strings.HasPrefix(path, "/v1/audio/translations"):
return c.Request.FormValue("model"), nil
case strings.HasPrefix(path, "/v1/engines") && strings.HasSuffix(path, "/embeddings"):
// /engines/:model/embeddings
return c.Param("model"), nil
default:
var modelRequest ModelRequest
err := common.UnmarshalBodyReusable(c.Request, &modelRequest)
if err != nil {
return "", fmt.Errorf("get request model failed: %w", err)
}
return modelRequest.Model, nil
}
}
+299 -224
View File
@@ -9,22 +9,23 @@ import (
"slices"
"sort"
"sync"
"sync/atomic"
"time"
json "github.com/json-iterator/go"
"github.com/maruel/natural"
"github.com/redis/go-redis/v9"
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/common/config"
"github.com/labring/sealos/service/aiproxy/common/conv"
"github.com/maruel/natural"
"github.com/redis/go-redis/v9"
log "github.com/sirupsen/logrus"
)
const (
SyncFrequency = time.Minute * 3
TokenCacheKey = "token:%s"
GroupCacheKey = "group:%s"
SyncFrequency = time.Minute * 3
TokenCacheKey = "token:%s"
GroupCacheKey = "group:%s"
GroupModelTPMKey = "group:%s:model_tpm"
)
var (
@@ -44,6 +45,11 @@ func (r redisStringSlice) MarshalBinary() ([]byte, error) {
type redisTime time.Time
var (
_ redis.Scanner = (*redisTime)(nil)
_ encoding.BinaryMarshaler = (*redisTime)(nil)
)
func (t *redisTime) ScanRedis(value string) error {
return (*time.Time)(t).UnmarshalBinary(conv.StringToBytes(value))
}
@@ -88,13 +94,13 @@ func CacheDeleteToken(key string) error {
}
//nolint:gosec
func CacheSetToken(token *Token) error {
func CacheSetToken(token *TokenCache) error {
if !common.RedisEnabled {
return nil
}
key := fmt.Sprintf(TokenCacheKey, token.Key)
pipe := common.RDB.Pipeline()
pipe.HSet(context.Background(), key, token.ToTokenCache())
pipe.HSet(context.Background(), key, token)
expireTime := SyncFrequency + time.Duration(rand.Int64N(60)-30)*time.Second
pipe.Expire(context.Background(), key, expireTime)
_, err := pipe.Exec(context.Background())
@@ -125,48 +131,27 @@ func CacheGetTokenByKey(key string) (*TokenCache, error) {
return nil, err
}
if err := CacheSetToken(token); err != nil {
tc := token.ToTokenCache()
if err := CacheSetToken(tc); err != nil {
log.Error("redis set token error: " + err.Error())
}
return token.ToTokenCache(), nil
return tc, nil
}
var updateTokenUsedAmountScript = redis.NewScript(`
if redis.call("HExists", KEYS[1], "used_amount") then
redis.call("HSet", KEYS[1], "used_amount", ARGV[1])
end
return redis.status_reply("ok")
`)
var updateTokenUsedAmountOnlyIncreaseScript = redis.NewScript(`
local used_amount = redis.call("HGet", KEYS[1], "used_amount")
local used_amount = redis.call("HGet", KEYS[1], "ua")
if used_amount == false then
return redis.status_reply("ok")
end
if ARGV[1] < used_amount then
return redis.status_reply("ok")
end
redis.call("HSet", KEYS[1], "used_amount", ARGV[1])
redis.call("HSet", KEYS[1], "ua", ARGV[1])
return redis.status_reply("ok")
`)
var increaseTokenUsedAmountScript = redis.NewScript(`
local used_amount = redis.call("HGet", KEYS[1], "used_amount")
if used_amount == false then
return redis.status_reply("ok")
end
redis.call("HSet", KEYS[1], "used_amount", used_amount + ARGV[1])
return redis.status_reply("ok")
`)
func CacheUpdateTokenUsedAmount(key string, amount float64) error {
if !common.RedisEnabled {
return nil
}
return updateTokenUsedAmountScript.Run(context.Background(), common.RDB, []string{fmt.Sprintf(TokenCacheKey, key)}, amount).Err()
}
func CacheUpdateTokenUsedAmountOnlyIncrease(key string, amount float64) error {
if !common.RedisEnabled {
return nil
@@ -174,24 +159,68 @@ func CacheUpdateTokenUsedAmountOnlyIncrease(key string, amount float64) error {
return updateTokenUsedAmountOnlyIncreaseScript.Run(context.Background(), common.RDB, []string{fmt.Sprintf(TokenCacheKey, key)}, amount).Err()
}
func CacheIncreaseTokenUsedAmount(key string, amount float64) error {
var updateTokenNameScript = redis.NewScript(`
if redis.call("HExists", KEYS[1], "n") then
redis.call("HSet", KEYS[1], "n", ARGV[1])
end
return redis.status_reply("ok")
`)
func CacheUpdateTokenName(key string, name string) error {
if !common.RedisEnabled {
return nil
}
return increaseTokenUsedAmountScript.Run(context.Background(), common.RDB, []string{fmt.Sprintf(TokenCacheKey, key)}, amount).Err()
return updateTokenNameScript.Run(context.Background(), common.RDB, []string{fmt.Sprintf(TokenCacheKey, key)}, name).Err()
}
var updateTokenStatusScript = redis.NewScript(`
if redis.call("HExists", KEYS[1], "st") then
redis.call("HSet", KEYS[1], "st", ARGV[1])
end
return redis.status_reply("ok")
`)
func CacheUpdateTokenStatus(key string, status int) error {
if !common.RedisEnabled {
return nil
}
return updateTokenStatusScript.Run(context.Background(), common.RDB, []string{fmt.Sprintf(TokenCacheKey, key)}, status).Err()
}
type redisMapStringInt64 map[string]int64
var (
_ redis.Scanner = (*redisMapStringInt64)(nil)
_ encoding.BinaryMarshaler = (*redisMapStringInt64)(nil)
)
func (r *redisMapStringInt64) ScanRedis(value string) error {
return json.Unmarshal(conv.StringToBytes(value), r)
}
func (r redisMapStringInt64) MarshalBinary() ([]byte, error) {
return json.Marshal(r)
}
type GroupCache struct {
ID string `json:"-" redis:"-"`
Status int `json:"status" redis:"st"`
QPM int64 `json:"qpm" redis:"q"`
ID string `json:"-" redis:"-"`
Status int `json:"status" redis:"st"`
UsedAmount float64 `json:"used_amount" redis:"ua"`
RPMRatio float64 `json:"rpm_ratio" redis:"rpm_r"`
RPM redisMapStringInt64 `json:"rpm" redis:"rpm"`
TPMRatio float64 `json:"tpm_ratio" redis:"tpm_r"`
TPM redisMapStringInt64 `json:"tpm" redis:"tpm"`
}
func (g *Group) ToGroupCache() *GroupCache {
return &GroupCache{
ID: g.ID,
Status: g.Status,
QPM: g.QPM,
ID: g.ID,
Status: g.Status,
UsedAmount: g.UsedAmount,
RPMRatio: g.RPMRatio,
RPM: g.RPM,
TPMRatio: g.TPMRatio,
TPM: g.TPM,
}
}
@@ -202,23 +231,73 @@ func CacheDeleteGroup(id string) error {
return common.RedisDel(fmt.Sprintf(GroupCacheKey, id))
}
var updateGroupQPMScript = redis.NewScript(`
if redis.call("HExists", KEYS[1], "qpm") then
redis.call("HSet", KEYS[1], "qpm", ARGV[1])
var updateGroupRPMRatioScript = redis.NewScript(`
if redis.call("HExists", KEYS[1], "rpm_r") then
redis.call("HSet", KEYS[1], "rpm_r", ARGV[1])
end
return redis.status_reply("ok")
`)
func CacheUpdateGroupQPM(id string, qpm int64) error {
func CacheUpdateGroupRPMRatio(id string, rpmRatio float64) error {
if !common.RedisEnabled {
return nil
}
return updateGroupQPMScript.Run(context.Background(), common.RDB, []string{fmt.Sprintf(GroupCacheKey, id)}, qpm).Err()
return updateGroupRPMRatioScript.Run(context.Background(), common.RDB, []string{fmt.Sprintf(GroupCacheKey, id)}, rpmRatio).Err()
}
var updateGroupRPMScript = redis.NewScript(`
if redis.call("HExists", KEYS[1], "rpm") then
redis.call("HSet", KEYS[1], "rpm", ARGV[1])
end
return redis.status_reply("ok")
`)
func CacheUpdateGroupRPM(id string, rpm map[string]int64) error {
if !common.RedisEnabled {
return nil
}
jsonRPM, err := json.Marshal(rpm)
if err != nil {
return err
}
return updateGroupRPMScript.Run(context.Background(), common.RDB, []string{fmt.Sprintf(GroupCacheKey, id)}, conv.BytesToString(jsonRPM)).Err()
}
var updateGroupTPMRatioScript = redis.NewScript(`
if redis.call("HExists", KEYS[1], "tpm_r") then
redis.call("HSet", KEYS[1], "tpm_r", ARGV[1])
end
return redis.status_reply("ok")
`)
func CacheUpdateGroupTPMRatio(id string, tpmRatio float64) error {
if !common.RedisEnabled {
return nil
}
return updateGroupTPMRatioScript.Run(context.Background(), common.RDB, []string{fmt.Sprintf(GroupCacheKey, id)}, tpmRatio).Err()
}
var updateGroupTPMScript = redis.NewScript(`
if redis.call("HExists", KEYS[1], "tpm") then
redis.call("HSet", KEYS[1], "tpm", ARGV[1])
end
return redis.status_reply("ok")
`)
func CacheUpdateGroupTPM(id string, tpm map[string]int64) error {
if !common.RedisEnabled {
return nil
}
jsonTPM, err := json.Marshal(tpm)
if err != nil {
return err
}
return updateGroupTPMScript.Run(context.Background(), common.RDB, []string{fmt.Sprintf(GroupCacheKey, id)}, conv.BytesToString(jsonTPM)).Err()
}
var updateGroupStatusScript = redis.NewScript(`
if redis.call("HExists", KEYS[1], "status") then
redis.call("HSet", KEYS[1], "status", ARGV[1])
if redis.call("HExists", KEYS[1], "st") then
redis.call("HSet", KEYS[1], "st", ARGV[1])
end
return redis.status_reply("ok")
`)
@@ -231,13 +310,13 @@ func CacheUpdateGroupStatus(id string, status int) error {
}
//nolint:gosec
func CacheSetGroup(group *Group) error {
func CacheSetGroup(group *GroupCache) error {
if !common.RedisEnabled {
return nil
}
key := fmt.Sprintf(GroupCacheKey, group.ID)
pipe := common.RDB.Pipeline()
pipe.HSet(context.Background(), key, group.ToGroupCache())
pipe.HSet(context.Background(), key, group)
expireTime := SyncFrequency + time.Duration(rand.Int64N(60)-30)*time.Second
pipe.Expire(context.Background(), key, expireTime)
_, err := pipe.Exec(context.Background())
@@ -268,89 +347,103 @@ func CacheGetGroup(id string) (*GroupCache, error) {
return nil, err
}
if err := CacheSetGroup(group); err != nil {
gc := group.ToGroupCache()
if err := CacheSetGroup(gc); err != nil {
log.Error("redis set group error: " + err.Error())
}
return group.ToGroupCache(), nil
return gc, nil
}
var (
enabledChannels []*Channel
allChannels []*Channel
enabledModel2channels map[string][]*Channel
enabledModels []string
enabledModelConfigs []*ModelConfig
enabledChannelType2ModelConfigs map[int][]*ModelConfig
enabledChannelID2channel map[int]*Channel
allChannelID2channel map[int]*Channel
channelSyncLock sync.RWMutex
)
var updateGroupUsedAmountOnlyIncreaseScript = redis.NewScript(`
local used_amount = redis.call("HGet", KEYS[1], "ua")
if used_amount == false then
return redis.status_reply("ok")
end
if ARGV[1] < used_amount then
return redis.status_reply("ok")
end
redis.call("HSet", KEYS[1], "ua", ARGV[1])
return redis.status_reply("ok")
`)
func CacheGetAllChannels() []*Channel {
channelSyncLock.RLock()
defer channelSyncLock.RUnlock()
return allChannels
func CacheUpdateGroupUsedAmountOnlyIncrease(id string, amount float64) error {
if !common.RedisEnabled {
return nil
}
return updateGroupUsedAmountOnlyIncreaseScript.Run(context.Background(), common.RDB, []string{fmt.Sprintf(GroupCacheKey, id)}, amount).Err()
}
func CacheGetAllChannelByID(id int) (*Channel, bool) {
channelSyncLock.RLock()
defer channelSyncLock.RUnlock()
channel, ok := allChannelID2channel[id]
return channel, ok
//nolint:gosec
func CacheGetGroupModelTPM(id string, model string) (int64, error) {
if !common.RedisEnabled {
return GetGroupModelTPM(id, model)
}
cacheKey := fmt.Sprintf(GroupModelTPMKey, id)
tpm, err := common.RDB.HGet(context.Background(), cacheKey, model).Int64()
if err == nil {
return tpm, nil
} else if !errors.Is(err, redis.Nil) {
log.Errorf("get group model tpm (%s:%s) from redis error: %s", id, model, err.Error())
}
tpm, err = GetGroupModelTPM(id, model)
if err != nil {
return 0, err
}
pipe := common.RDB.Pipeline()
pipe.HSet(context.Background(), cacheKey, model, tpm)
// 2-5 seconds
pipe.Expire(context.Background(), cacheKey, 2*time.Second+time.Duration(rand.Int64N(3))*time.Second)
_, err = pipe.Exec(context.Background())
if err != nil {
log.Errorf("set group model tpm (%s:%s) to redis error: %s", id, model, err.Error())
}
return tpm, nil
}
// GetEnabledModel2Channels returns a map of model name to enabled channels
func GetEnabledModel2Channels() map[string][]*Channel {
channelSyncLock.RLock()
defer channelSyncLock.RUnlock()
return enabledModel2channels
//nolint:revive
type ModelConfigCache interface {
GetModelConfig(model string) (*ModelConfig, bool)
}
// CacheGetEnabledModels returns a list of enabled model names
func CacheGetEnabledModels() []string {
channelSyncLock.RLock()
defer channelSyncLock.RUnlock()
return enabledModels
// read-only cache
//
//nolint:revive
type ModelCaches struct {
ModelConfig ModelConfigCache
EnabledModel2channels map[string][]*Channel
EnabledModels []string
EnabledModelsMap map[string]struct{}
EnabledModelConfigs []*ModelConfig
EnabledModelConfigsMap map[string]*ModelConfig
EnabledChannelType2ModelConfigs map[int][]*ModelConfig
EnabledChannelID2channel map[int]*Channel
}
// CacheGetEnabledChannelType2ModelConfigs returns a map of channel type to enabled model configs
func CacheGetEnabledChannelType2ModelConfigs() map[int][]*ModelConfig {
channelSyncLock.RLock()
defer channelSyncLock.RUnlock()
return enabledChannelType2ModelConfigs
var modelCaches atomic.Pointer[ModelCaches]
func init() {
modelCaches.Store(new(ModelCaches))
}
// CacheGetEnabledModelConfigs returns a list of enabled model configs
func CacheGetEnabledModelConfigs() []*ModelConfig {
channelSyncLock.RLock()
defer channelSyncLock.RUnlock()
return enabledModelConfigs
func LoadModelCaches() *ModelCaches {
return modelCaches.Load()
}
func CacheGetEnabledChannels() []*Channel {
channelSyncLock.RLock()
defer channelSyncLock.RUnlock()
return enabledChannels
}
func CacheGetEnabledChannelByID(id int) (*Channel, bool) {
channelSyncLock.RLock()
defer channelSyncLock.RUnlock()
channel, ok := enabledChannelID2channel[id]
return channel, ok
}
// InitChannelCache initializes the channel cache from database
func InitChannelCache() error {
// Load enabled newEnabledChannels from database
newEnabledChannels, err := LoadEnabledChannels()
// InitModelConfigAndChannelCache initializes the channel cache from database
func InitModelConfigAndChannelCache() error {
modelConfig, err := initializeModelConfigCache()
if err != nil {
return err
}
// Load all channels from database
newAllChannels, err := LoadChannels()
// Load enabled newEnabledChannels from database
newEnabledChannels, err := LoadEnabledChannels()
if err != nil {
return err
}
@@ -359,7 +452,6 @@ func InitChannelCache() error {
newEnabledChannelID2channel := buildChannelIDMap(newEnabledChannels)
// Build all channel ID to channel map
newAllChannelID2channel := buildChannelIDMap(newAllChannels)
// Build model to channels map
newEnabledModel2channels := buildModelToChannelsMap(newEnabledChannels)
@@ -368,29 +460,29 @@ func InitChannelCache() error {
sortChannelsByPriority(newEnabledModel2channels)
// Build channel type to model configs map
newEnabledChannelType2ModelConfigs := buildChannelTypeToModelConfigsMap(newEnabledChannels)
newEnabledChannelType2ModelConfigs := buildChannelTypeToModelConfigsMap(newEnabledChannels, modelConfig)
// Build enabled models and configs lists
newEnabledModels, newEnabledModelConfigs := buildEnabledModelsAndConfigs(newEnabledChannelType2ModelConfigs)
newEnabledModels, newEnabledModelsMap, newEnabledModelConfigs, newEnabledModelConfigsMap := buildEnabledModelsAndConfigs(newEnabledChannelType2ModelConfigs)
// Update global cache atomically
updateGlobalCache(
newEnabledChannels,
newAllChannels,
newEnabledModel2channels,
newEnabledModels,
newEnabledModelConfigs,
newEnabledChannelID2channel,
newEnabledChannelType2ModelConfigs,
newAllChannelID2channel,
)
modelCaches.Store(&ModelCaches{
ModelConfig: modelConfig,
EnabledModel2channels: newEnabledModel2channels,
EnabledModels: newEnabledModels,
EnabledModelsMap: newEnabledModelsMap,
EnabledModelConfigs: newEnabledModelConfigs,
EnabledModelConfigsMap: newEnabledModelConfigsMap,
EnabledChannelType2ModelConfigs: newEnabledChannelType2ModelConfigs,
EnabledChannelID2channel: newEnabledChannelID2channel,
})
return nil
}
func LoadEnabledChannels() ([]*Channel, error) {
var channels []*Channel
err := DB.Where("status = ?", ChannelStatusEnabled).Find(&channels).Error
err := DB.Where("status = ? or status = ?", ChannelStatusEnabled, ChannelStatusFail).Find(&channels).Error
if err != nil {
return nil, err
}
@@ -431,13 +523,54 @@ func LoadChannelByID(id int) (*Channel, error) {
return &channel, nil
}
var _ ModelConfigCache = (*modelConfigMapCache)(nil)
type modelConfigMapCache struct {
modelConfigMap map[string]*ModelConfig
}
func (m *modelConfigMapCache) GetModelConfig(model string) (*ModelConfig, bool) {
config, ok := m.modelConfigMap[model]
return config, ok
}
var _ ModelConfigCache = (*disabledModelConfigCache)(nil)
type disabledModelConfigCache struct {
modelConfigs ModelConfigCache
}
func (d *disabledModelConfigCache) GetModelConfig(model string) (*ModelConfig, bool) {
if config, ok := d.modelConfigs.GetModelConfig(model); ok {
return config, true
}
return NewDefaultModelConfig(model), true
}
func initializeModelConfigCache() (ModelConfigCache, error) {
modelConfigs, err := GetAllModelConfigs()
if err != nil {
return nil, err
}
newModelConfigMap := make(map[string]*ModelConfig)
for _, modelConfig := range modelConfigs {
newModelConfigMap[modelConfig.Model] = modelConfig
}
configs := &modelConfigMapCache{modelConfigMap: newModelConfigMap}
if config.GetDisableModelConfig() {
return &disabledModelConfigCache{modelConfigs: configs}, nil
}
return configs, nil
}
func initializeChannelModels(channel *Channel) {
if len(channel.Models) == 0 {
channel.Models = config.GetDefaultChannelModels()[channel.Type]
return
}
findedModels, missingModels, err := CheckModelConfig(channel.Models)
findedModels, missingModels, err := GetModelConfigWithModels(channel.Models)
if err != nil {
return
}
@@ -477,12 +610,12 @@ func buildModelToChannelsMap(channels []*Channel) map[string][]*Channel {
func sortChannelsByPriority(modelMap map[string][]*Channel) {
for _, channels := range modelMap {
sort.Slice(channels, func(i, j int) bool {
return channels[i].Priority > channels[j].Priority
return channels[i].GetPriority() > channels[j].GetPriority()
})
}
}
func buildChannelTypeToModelConfigsMap(channels []*Channel) map[int][]*ModelConfig {
func buildChannelTypeToModelConfigsMap(channels []*Channel, modelConfigMap ModelConfigCache) map[int][]*ModelConfig {
typeMap := make(map[int][]*ModelConfig)
for _, channel := range channels {
@@ -492,7 +625,7 @@ func buildChannelTypeToModelConfigsMap(channels []*Channel) map[int][]*ModelConf
configs := typeMap[channel.Type]
for _, model := range channel.Models {
if config, ok := CacheGetModelConfig(model); ok {
if config, ok := modelConfigMap.GetModelConfig(model); ok {
configs = append(configs, config)
}
}
@@ -508,10 +641,11 @@ func buildChannelTypeToModelConfigsMap(channels []*Channel) map[int][]*ModelConf
return typeMap
}
func buildEnabledModelsAndConfigs(typeMap map[int][]*ModelConfig) ([]string, []*ModelConfig) {
func buildEnabledModelsAndConfigs(typeMap map[int][]*ModelConfig) ([]string, map[string]struct{}, []*ModelConfig, map[string]*ModelConfig) {
models := make([]string, 0)
configs := make([]*ModelConfig, 0)
appended := make(map[string]struct{})
modelConfigsMap := make(map[string]*ModelConfig)
for _, modelConfigs := range typeMap {
for _, config := range modelConfigs {
@@ -521,13 +655,14 @@ func buildEnabledModelsAndConfigs(typeMap map[int][]*ModelConfig) ([]string, []*
models = append(models, config.Model)
configs = append(configs, config)
appended[config.Model] = struct{}{}
modelConfigsMap[config.Model] = config
}
}
slices.Sort(models)
slices.SortStableFunc(configs, SortModelConfigsFunc)
return models, configs
return models, appended, configs, modelConfigsMap
}
func SortModelConfigsFunc(i, j *ModelConfig) int {
@@ -552,29 +687,7 @@ func SortModelConfigsFunc(i, j *ModelConfig) int {
return 1
}
func updateGlobalCache(
newEnabledChannels []*Channel,
newAllChannels []*Channel,
newEnabledModel2channels map[string][]*Channel,
newEnabledModels []string,
newEnabledModelConfigs []*ModelConfig,
newEnabledChannelID2channel map[int]*Channel,
newEnabledChannelType2ModelConfigs map[int][]*ModelConfig,
newAllChannelID2channel map[int]*Channel,
) {
channelSyncLock.Lock()
defer channelSyncLock.Unlock()
enabledChannels = newEnabledChannels
allChannels = newAllChannels
enabledModel2channels = newEnabledModel2channels
enabledModels = newEnabledModels
enabledModelConfigs = newEnabledModelConfigs
enabledChannelID2channel = newEnabledChannelID2channel
enabledChannelType2ModelConfigs = newEnabledChannelType2ModelConfigs
allChannelID2channel = newAllChannelID2channel
}
func SyncChannelCache(ctx context.Context, wg *sync.WaitGroup, frequency time.Duration) {
func SyncModelConfigAndChannelCache(ctx context.Context, wg *sync.WaitGroup, frequency time.Duration) {
defer wg.Done()
ticker := time.NewTicker(frequency)
@@ -584,7 +697,7 @@ func SyncChannelCache(ctx context.Context, wg *sync.WaitGroup, frequency time.Du
case <-ctx.Done():
return
case <-ticker.C:
err := InitChannelCache()
err := InitModelConfigAndChannelCache()
if err != nil {
log.Error("failed to sync channels: " + err.Error())
continue
@@ -593,11 +706,35 @@ func SyncChannelCache(ctx context.Context, wg *sync.WaitGroup, frequency time.Du
}
}
func filterChannels(channels []*Channel, ignoreChannel ...int) []*Channel {
filtered := make([]*Channel, 0)
for _, channel := range channels {
if channel.Status != ChannelStatusEnabled {
continue
}
if slices.Contains(ignoreChannel, channel.ID) {
continue
}
filtered = append(filtered, channel)
}
return filtered
}
var (
ErrChannelsNotFound = errors.New("channels not found")
ErrChannelsExhausted = errors.New("channels exhausted")
)
//nolint:gosec
func CacheGetRandomSatisfiedChannel(model string) (*Channel, error) {
channels := GetEnabledModel2Channels()[model]
func (c *ModelCaches) GetRandomSatisfiedChannel(model string, ignoreChannel ...int) (*Channel, error) {
_channels := c.EnabledModel2channels[model]
if len(_channels) == 0 {
return nil, ErrChannelsNotFound
}
channels := filterChannels(_channels, ignoreChannel...)
if len(channels) == 0 {
return nil, errors.New("model not found")
return nil, ErrChannelsExhausted
}
if len(channels) == 1 {
@@ -606,7 +743,7 @@ func CacheGetRandomSatisfiedChannel(model string) (*Channel, error) {
var totalWeight int32
for _, ch := range channels {
totalWeight += ch.Priority
totalWeight += ch.GetPriority()
}
if totalWeight == 0 {
@@ -615,7 +752,7 @@ func CacheGetRandomSatisfiedChannel(model string) (*Channel, error) {
r := rand.Int32N(totalWeight)
for _, ch := range channels {
r -= ch.Priority
r -= ch.GetPriority()
if r < 0 {
return ch, nil
}
@@ -623,65 +760,3 @@ func CacheGetRandomSatisfiedChannel(model string) (*Channel, error) {
return channels[rand.IntN(len(channels))], nil
}
var (
modelConfigSyncLock sync.RWMutex
modelConfigMap map[string]*ModelConfig
)
func InitModelConfigCache() error {
modelConfigs, err := GetAllModelConfigs()
if err != nil {
return err
}
newModelConfigMap := make(map[string]*ModelConfig)
for _, modelConfig := range modelConfigs {
newModelConfigMap[modelConfig.Model] = modelConfig
}
modelConfigSyncLock.Lock()
modelConfigMap = newModelConfigMap
modelConfigSyncLock.Unlock()
return nil
}
func SyncModelConfigCache(ctx context.Context, wg *sync.WaitGroup, frequency time.Duration) {
defer wg.Done()
ticker := time.NewTicker(frequency)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
err := InitModelConfigCache()
if err != nil {
log.Error("failed to sync model configs: " + err.Error())
}
}
}
}
func CacheGetModelConfig(model string) (*ModelConfig, bool) {
modelConfigSyncLock.RLock()
defer modelConfigSyncLock.RUnlock()
modelConfig, ok := modelConfigMap[model]
return modelConfig, ok
}
func CacheCheckModelConfig(models []string) ([]string, []string) {
if len(models) == 0 {
return models, nil
}
founded := make([]string, 0)
missing := make([]string, 0)
for _, model := range models {
if _, ok := modelConfigMap[model]; ok {
founded = append(founded, model)
} else {
missing = append(missing, model)
}
}
return founded, missing
}
+85 -95
View File
@@ -2,13 +2,13 @@ package model
import (
"fmt"
"slices"
"strings"
"time"
json "github.com/json-iterator/go"
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/common/helper"
"github.com/labring/sealos/service/aiproxy/common/config"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
@@ -18,20 +18,22 @@ const (
)
const (
ChannelStatusUnknown = 0
ChannelStatusEnabled = 1 // don't use 0, 0 is the default value!
ChannelStatusManuallyDisabled = 2 // also don't use 0
ChannelStatusAutoDisabled = 3
ChannelStatusUnknown = 0
ChannelStatusEnabled = 1 // don't use 0, 0 is the default value!
ChannelStatusDisabled = 2 // also don't use 0
ChannelStatusFail = 3
)
type ChannelConfig struct {
SplitThink bool `json:"split_think"`
}
type Channel struct {
CreatedAt time.Time `gorm:"index" json:"created_at"`
AccessedAt time.Time `json:"accessed_at"`
LastTestErrorAt time.Time `json:"last_test_error_at"`
ChannelTests []*ChannelTest `gorm:"foreignKey:ChannelID;references:ID" json:"channel_tests"`
ChannelTests []*ChannelTest `gorm:"foreignKey:ChannelID;references:ID" json:"channel_tests,omitempty"`
BalanceUpdatedAt time.Time `json:"balance_updated_at"`
ModelMapping map[string]string `gorm:"serializer:fastjson;type:text" json:"model_mapping"`
Config ChannelConfig `gorm:"serializer:fastjson;type:text" json:"config"`
Key string `gorm:"type:text;index" json:"key"`
Name string `gorm:"index" json:"name"`
BaseURL string `gorm:"index" json:"base_url"`
@@ -43,37 +45,30 @@ type Channel struct {
Status int `gorm:"default:1;index" json:"status"`
Type int `gorm:"default:0;index" json:"type"`
Priority int32 `json:"priority"`
Config *ChannelConfig `gorm:"serializer:fastjson;type:text" json:"config,omitempty"`
}
func (c *Channel) BeforeDelete(tx *gorm.DB) (err error) {
return tx.Model(&ChannelTest{}).Where("channel_id = ?", c.ID).Delete(&ChannelTest{}).Error
}
// check model config exist
func (c *Channel) BeforeSave(tx *gorm.DB) (err error) {
if len(c.Models) == 0 {
return nil
const (
DefaultPriority = 100
)
func (c *Channel) GetPriority() int32 {
if c.Priority == 0 {
return DefaultPriority
}
_, missingModels, err := checkModelConfig(tx, c.Models)
if err != nil {
return err
}
if len(missingModels) > 0 {
return fmt.Errorf("model config not found: %v", missingModels)
}
return nil
return c.Priority
}
func CheckModelConfig(models []string) ([]string, []string, error) {
return checkModelConfig(DB, models)
}
func checkModelConfig(tx *gorm.DB, models []string) ([]string, []string, error) {
if len(models) == 0 {
func GetModelConfigWithModels(models []string) ([]string, []string, error) {
if len(models) == 0 || config.GetDisableModelConfig() {
return models, nil, nil
}
where := tx.Model(&ModelConfig{}).Where("model IN ?", models)
where := DB.Model(&ModelConfig{}).Where("model IN ?", models)
var count int64
if err := where.Count(&count).Error; err != nil {
return nil, nil, err
@@ -108,18 +103,28 @@ func checkModelConfig(tx *gorm.DB, models []string) ([]string, []string, error)
return foundModels, nil, nil
}
func CheckModelConfigExist(models []string) error {
_, missingModels, err := GetModelConfigWithModels(models)
if err != nil {
return err
}
if len(missingModels) > 0 {
slices.Sort(missingModels)
return fmt.Errorf("model config not found: %v", missingModels)
}
return nil
}
func (c *Channel) MarshalJSON() ([]byte, error) {
type Alias Channel
return json.Marshal(&struct {
*Alias
CreatedAt int64 `json:"created_at"`
AccessedAt int64 `json:"accessed_at"`
BalanceUpdatedAt int64 `json:"balance_updated_at"`
LastTestErrorAt int64 `json:"last_test_error_at"`
}{
Alias: (*Alias)(c),
CreatedAt: c.CreatedAt.UnixMilli(),
AccessedAt: c.AccessedAt.UnixMilli(),
BalanceUpdatedAt: c.BalanceUpdatedAt.UnixMilli(),
LastTestErrorAt: c.LastTestErrorAt.UnixMilli(),
})
@@ -129,7 +134,7 @@ func (c *Channel) MarshalJSON() ([]byte, error) {
func getChannelOrder(order string) string {
prefix, suffix, _ := strings.Cut(order, "-")
switch prefix {
case "name", "type", "created_at", "accessed_at", "status", "test_at", "balance_updated_at", "used_amount", "request_count", "priority", "id":
case "name", "type", "created_at", "status", "test_at", "balance_updated_at", "used_amount", "request_count", "priority", "id":
switch suffix {
case "asc":
return prefix + " asc"
@@ -141,34 +146,14 @@ func getChannelOrder(order string) string {
}
}
type ChannelConfig struct {
Region string `json:"region,omitempty"`
SK string `json:"sk,omitempty"`
AK string `json:"ak,omitempty"`
UserID string `json:"user_id,omitempty"`
APIVersion string `json:"api_version,omitempty"`
Plugin string `json:"plugin,omitempty"`
VertexAIProjectID string `json:"vertex_ai_project_id,omitempty"`
VertexAIADC string `json:"vertex_ai_adc,omitempty"`
}
func GetAllChannels(onlyDisabled bool, omitKey bool) (channels []*Channel, err error) {
func GetAllChannels() (channels []*Channel, err error) {
tx := DB.Model(&Channel{})
if onlyDisabled {
tx = tx.Where("status = ? or status = ?", ChannelStatusAutoDisabled, ChannelStatusManuallyDisabled)
}
if omitKey {
tx = tx.Omit("key")
}
err = tx.Order("id desc").Find(&channels).Error
return channels, err
}
func GetChannels(startIdx int, num int, onlyDisabled bool, omitKey bool, id int, name string, key string, channelType int, baseURL string, order string) (channels []*Channel, total int64, err error) {
func GetChannels(startIdx int, num int, id int, name string, key string, channelType int, baseURL string, order string) (channels []*Channel, total int64, err error) {
tx := DB.Model(&Channel{})
if onlyDisabled {
tx = tx.Where("status = ? or status = ?", ChannelStatusAutoDisabled, ChannelStatusManuallyDisabled)
}
if id != 0 {
tx = tx.Where("id = ?", id)
}
@@ -188,9 +173,6 @@ func GetChannels(startIdx int, num int, onlyDisabled bool, omitKey bool, id int,
if err != nil {
return nil, 0, err
}
if omitKey {
tx = tx.Omit("key")
}
if total <= 0 {
return nil, 0, nil
}
@@ -198,11 +180,8 @@ func GetChannels(startIdx int, num int, onlyDisabled bool, omitKey bool, id int,
return channels, total, err
}
func SearchChannels(keyword string, startIdx int, num int, onlyDisabled bool, omitKey bool, id int, name string, key string, channelType int, baseURL string, order string) (channels []*Channel, total int64, err error) {
func SearchChannels(keyword string, startIdx int, num int, id int, name string, key string, channelType int, baseURL string, order string) (channels []*Channel, total int64, err error) {
tx := DB.Model(&Channel{})
if onlyDisabled {
tx = tx.Where("status = ? or status = ?", ChannelStatusAutoDisabled, ChannelStatusManuallyDisabled)
}
// Handle exact match conditions for non-zero values
if id != 0 {
@@ -228,7 +207,11 @@ func SearchChannels(keyword string, startIdx int, num int, onlyDisabled bool, om
if id == 0 {
conditions = append(conditions, "id = ?")
values = append(values, helper.String2Int(keyword))
values = append(values, String2Int(keyword))
}
if channelType == 0 {
conditions = append(conditions, "type = ?")
values = append(values, String2Int(keyword))
}
if name == "" {
if common.UsingPostgreSQL {
@@ -246,10 +229,6 @@ func SearchChannels(keyword string, startIdx int, num int, onlyDisabled bool, om
}
values = append(values, "%"+keyword+"%")
}
if channelType == 0 {
conditions = append(conditions, "type = ?")
values = append(values, helper.String2Int(keyword))
}
if baseURL == "" {
if common.UsingPostgreSQL {
conditions = append(conditions, "base_url ILIKE ?")
@@ -259,6 +238,13 @@ func SearchChannels(keyword string, startIdx int, num int, onlyDisabled bool, om
values = append(values, "%"+keyword+"%")
}
if common.UsingPostgreSQL {
conditions = append(conditions, "models ILIKE ?")
} else {
conditions = append(conditions, "models LIKE ?")
}
values = append(values, "%"+keyword+"%")
if len(conditions) > 0 {
tx = tx.Where(fmt.Sprintf("(%s)", strings.Join(conditions, " OR ")), values...)
}
@@ -268,9 +254,6 @@ func SearchChannels(keyword string, startIdx int, num int, onlyDisabled bool, om
if err != nil {
return nil, 0, err
}
if omitKey {
tx = tx.Omit("key")
}
if total <= 0 {
return nil, 0, nil
}
@@ -278,28 +261,32 @@ func SearchChannels(keyword string, startIdx int, num int, onlyDisabled bool, om
return channels, total, err
}
func GetChannelByID(id int, omitKey bool) (*Channel, error) {
func GetChannelByID(id int) (*Channel, error) {
channel := Channel{ID: id}
var err error
if omitKey {
err = DB.Omit("key").First(&channel, "id = ?", id).Error
} else {
err = DB.First(&channel, "id = ?", id).Error
}
err := DB.First(&channel, "id = ?", id).Error
return &channel, HandleNotFound(err, ErrChannelNotFound)
}
func BatchInsertChannels(channels []*Channel) error {
for _, channel := range channels {
if err := CheckModelConfigExist(channel.Models); err != nil {
return err
}
}
return DB.Transaction(func(tx *gorm.DB) error {
return tx.Create(&channels).Error
})
}
func UpdateChannel(channel *Channel) error {
if err := CheckModelConfigExist(channel.Models); err != nil {
return err
}
result := DB.
Model(channel).
Omit("accessed_at", "used_amount", "request_count", "created_at", "balance_updated_at", "balance").
Omit("used_amount", "request_count", "created_at", "balance_updated_at", "balance").
Clauses(clause.Returning{}).
Where("id = ?", channel.ID).
Updates(channel)
return HandleUpdateResult(result, ErrChannelNotFound)
}
@@ -346,10 +333,13 @@ func (c *Channel) UpdateModelTest(testAt time.Time, model, actualModel string, m
}
func (c *Channel) UpdateBalance(balance float64) error {
result := DB.Model(c).Select("balance_updated_at", "balance").Updates(Channel{
BalanceUpdatedAt: time.Now(),
Balance: balance,
})
result := DB.Model(&Channel{}).
Select("balance_updated_at", "balance").
Where("id = ?", c.ID).
Updates(Channel{
BalanceUpdatedAt: time.Now(),
Balance: balance,
})
return HandleUpdateResult(result, ErrChannelNotFound)
}
@@ -368,28 +358,28 @@ func DeleteChannelsByIDs(ids []int) error {
}
func UpdateChannelStatusByID(id int, status int) error {
result := DB.Model(&Channel{}).Where("id = ?", id).Update("status", status)
result := DB.Model(&Channel{}).
Where("id = ?", id).
Update("status", status)
return HandleUpdateResult(result, ErrChannelNotFound)
}
func DisableChannelByID(id int) error {
return UpdateChannelStatusByID(id, ChannelStatusAutoDisabled)
}
func EnableChannelByID(id int) error {
return UpdateChannelStatusByID(id, ChannelStatusEnabled)
}
func UpdateChannelUsedAmount(id int, amount float64, requestCount int) error {
result := DB.Model(&Channel{}).Where("id = ?", id).Updates(map[string]interface{}{
"used_amount": gorm.Expr("used_amount + ?", amount),
"request_count": gorm.Expr("request_count + ?", requestCount),
"accessed_at": time.Now(),
})
result := DB.Model(&Channel{}).
Where("id = ?", id).
Updates(map[string]interface{}{
"used_amount": gorm.Expr("used_amount + ?", amount),
"request_count": gorm.Expr("request_count + ?", requestCount),
})
return HandleUpdateResult(result, ErrChannelNotFound)
}
func DeleteDisabledChannel() error {
result := DB.Where("status = ? or status = ?", ChannelStatusAutoDisabled, ChannelStatusManuallyDisabled).Delete(&Channel{})
result := DB.Where("status = ?", ChannelStatusDisabled).Delete(&Channel{})
return HandleUpdateResult(result, ErrChannelNotFound)
}
func DeleteFailChannel() error {
result := DB.Where("status = ?", ChannelStatusFail).Delete(&Channel{})
return HandleUpdateResult(result, ErrChannelNotFound)
}
+140
View File
@@ -0,0 +1,140 @@
package model
import "reflect"
//nolint:revive
type ModelConfigKey string
const (
ModelConfigMaxContextTokensKey ModelConfigKey = "max_context_tokens"
ModelConfigMaxInputTokensKey ModelConfigKey = "max_input_tokens"
ModelConfigMaxOutputTokensKey ModelConfigKey = "max_output_tokens"
ModelConfigVisionKey ModelConfigKey = "vision"
ModelConfigToolChoiceKey ModelConfigKey = "tool_choice"
ModelConfigSupportFormatsKey ModelConfigKey = "support_formats"
ModelConfigSupportVoicesKey ModelConfigKey = "support_voices"
)
//nolint:revive
type ModelConfigOption func(config map[ModelConfigKey]any)
func WithModelConfigMaxContextTokens(maxContextTokens int) ModelConfigOption {
return func(config map[ModelConfigKey]any) {
config[ModelConfigMaxContextTokensKey] = maxContextTokens
}
}
func WithModelConfigMaxInputTokens(maxInputTokens int) ModelConfigOption {
return func(config map[ModelConfigKey]any) {
config[ModelConfigMaxInputTokensKey] = maxInputTokens
}
}
func WithModelConfigMaxOutputTokens(maxOutputTokens int) ModelConfigOption {
return func(config map[ModelConfigKey]any) {
config[ModelConfigMaxOutputTokensKey] = maxOutputTokens
}
}
func WithModelConfigVision(vision bool) ModelConfigOption {
return func(config map[ModelConfigKey]any) {
config[ModelConfigVisionKey] = vision
}
}
func WithModelConfigToolChoice(toolChoice bool) ModelConfigOption {
return func(config map[ModelConfigKey]any) {
config[ModelConfigToolChoiceKey] = toolChoice
}
}
func WithModelConfigSupportFormats(supportFormats []string) ModelConfigOption {
return func(config map[ModelConfigKey]any) {
config[ModelConfigSupportFormatsKey] = supportFormats
}
}
func WithModelConfigSupportVoices(supportVoices []string) ModelConfigOption {
return func(config map[ModelConfigKey]any) {
config[ModelConfigSupportVoicesKey] = supportVoices
}
}
func NewModelConfig(opts ...ModelConfigOption) map[ModelConfigKey]any {
config := make(map[ModelConfigKey]any)
for _, opt := range opts {
opt(config)
}
return config
}
func GetModelConfigInt(config map[ModelConfigKey]any, key ModelConfigKey) (int, bool) {
if v, ok := config[key]; ok {
value := reflect.ValueOf(v)
if value.CanInt() {
return int(value.Int()), true
}
if value.CanFloat() {
return int(value.Float()), true
}
}
return 0, false
}
func GetModelConfigUint(config map[ModelConfigKey]any, key ModelConfigKey) (uint64, bool) {
if v, ok := config[key]; ok {
value := reflect.ValueOf(v)
if value.CanUint() {
return value.Uint(), true
}
if value.CanFloat() {
return uint64(value.Float()), true
}
}
return 0, false
}
func GetModelConfigFloat(config map[ModelConfigKey]any, key ModelConfigKey) (float64, bool) {
if v, ok := config[key]; ok {
value := reflect.ValueOf(v)
if value.CanFloat() {
return value.Float(), true
}
if value.CanInt() {
return float64(value.Int()), true
}
if value.CanUint() {
return float64(value.Uint()), true
}
}
return 0, false
}
func GetModelConfigStringSlice(config map[ModelConfigKey]any, key ModelConfigKey) ([]string, bool) {
v, ok := config[key]
if !ok {
return nil, false
}
if slice, ok := v.([]string); ok {
return slice, true
}
if slice, ok := v.([]any); ok {
result := make([]string, len(slice))
for i, v := range slice {
if s, ok := v.(string); ok {
result[i] = s
continue
}
return nil, false
}
return result, true
}
return nil, false
}
func GetModelConfigBool(config map[ModelConfigKey]any, key ModelConfigKey) (bool, bool) {
if v, ok := config[key].(bool); ok {
return v, true
}
return false, false
}
+1 -2
View File
@@ -7,7 +7,6 @@ import (
json "github.com/json-iterator/go"
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/common/helper"
)
type ConsumeError struct {
@@ -82,7 +81,7 @@ func SearchConsumeError(keyword string, requestID string, group string, tokenNam
if tokenID == 0 {
conditions = append(conditions, "token_id = ?")
values = append(values, helper.String2Int(keyword))
values = append(values, String2Int(keyword))
}
if requestID == "" {
if common.UsingPostgreSQL {
+87 -69
View File
@@ -2,12 +2,10 @@ package model
import (
"errors"
"fmt"
"strings"
"time"
json "github.com/json-iterator/go"
"github.com/labring/sealos/service/aiproxy/common"
log "github.com/sirupsen/logrus"
"gorm.io/gorm"
@@ -21,41 +19,31 @@ const (
const (
GroupStatusEnabled = 1 // don't use 0, 0 is the default value!
GroupStatusDisabled = 2 // also don't use 0
GroupStatusInternal = 3
)
type Group struct {
CreatedAt time.Time `json:"created_at"`
AccessedAt time.Time `json:"accessed_at"`
ID string `gorm:"primaryKey" json:"id"`
Tokens []*Token `gorm:"foreignKey:GroupID" json:"-"`
Status int `gorm:"default:1;index" json:"status"`
UsedAmount float64 `gorm:"index" json:"used_amount"`
QPM int64 `gorm:"index" json:"qpm"`
RequestCount int `gorm:"index" json:"request_count"`
CreatedAt time.Time `json:"created_at"`
ID string `gorm:"primaryKey" json:"id"`
Tokens []*Token `gorm:"foreignKey:GroupID" json:"-"`
Status int `gorm:"default:1;index" json:"status"`
UsedAmount float64 `gorm:"index" json:"used_amount"`
RPMRatio float64 `gorm:"index" json:"rpm_ratio"`
RPM map[string]int64 `gorm:"serializer:fastjson" json:"rpm"`
TPMRatio float64 `gorm:"index" json:"tpm_ratio"`
TPM map[string]int64 `gorm:"serializer:fastjson" json:"tpm"`
RequestCount int `gorm:"index" json:"request_count"`
}
func (g *Group) BeforeDelete(tx *gorm.DB) (err error) {
return tx.Model(&Token{}).Where("group_id = ?", g.ID).Delete(&Token{}).Error
}
func (g *Group) MarshalJSON() ([]byte, error) {
type Alias Group
return json.Marshal(&struct {
*Alias
CreatedAt int64 `json:"created_at"`
AccessedAt int64 `json:"accessed_at"`
}{
Alias: (*Alias)(g),
CreatedAt: g.CreatedAt.UnixMilli(),
AccessedAt: g.AccessedAt.UnixMilli(),
})
}
//nolint:goconst
func getGroupOrder(order string) string {
prefix, suffix, _ := strings.Cut(order, "-")
switch prefix {
case "id", "request_count", "accessed_at", "status", "created_at", "used_amount":
case "id", "request_count", "status", "created_at", "used_amount":
switch suffix {
case "asc":
return prefix + " asc"
@@ -88,7 +76,7 @@ func GetGroups(startIdx int, num int, order string, onlyDisabled bool) (groups [
func GetGroupByID(id string) (*Group, error) {
if id == "" {
return nil, errors.New("id 为空!")
return nil, errors.New("group id is empty")
}
group := Group{ID: id}
err := DB.First(&group, "id = ?", id).Error
@@ -97,7 +85,7 @@ func GetGroupByID(id string) (*Group, error) {
func DeleteGroupByID(id string) (err error) {
if id == "" {
return errors.New("id 为空!")
return errors.New("group id is empty")
}
defer func() {
if err == nil {
@@ -143,40 +131,83 @@ func DeleteGroupsByIDs(ids []string) (err error) {
})
}
func UpdateGroupUsedAmountAndRequestCount(id string, amount float64, count int) error {
result := DB.Model(&Group{}).Where("id = ?", id).Updates(map[string]interface{}{
"used_amount": gorm.Expr("used_amount + ?", amount),
"request_count": gorm.Expr("request_count + ?", count),
"accessed_at": time.Now(),
})
return HandleUpdateResult(result, ErrGroupNotFound)
}
func UpdateGroupUsedAmount(id string, amount float64) error {
result := DB.Model(&Group{}).Where("id = ?", id).Updates(map[string]interface{}{
"used_amount": gorm.Expr("used_amount + ?", amount),
"accessed_at": time.Now(),
})
return HandleUpdateResult(result, ErrGroupNotFound)
}
func UpdateGroupRequestCount(id string, count int) error {
result := DB.Model(&Group{}).Where("id = ?", id).Updates(map[string]interface{}{
"request_count": gorm.Expr("request_count + ?", count),
"accessed_at": time.Now(),
})
return HandleUpdateResult(result, ErrGroupNotFound)
}
func UpdateGroupQPM(id string, qpm int64) (err error) {
func UpdateGroupUsedAmountAndRequestCount(id string, amount float64, count int) (err error) {
group := &Group{ID: id}
defer func() {
if err == nil {
if err := CacheUpdateGroupQPM(id, qpm); err != nil {
log.Error("cache update group qpm failed: " + err.Error())
if amount > 0 && err == nil {
if err := CacheUpdateGroupUsedAmountOnlyIncrease(group.ID, group.UsedAmount); err != nil {
log.Error("update group used amount in cache failed: " + err.Error())
}
}
}()
result := DB.Model(&Group{}).Where("id = ?", id).Update("qpm", qpm)
result := DB.
Model(group).
Clauses(clause.Returning{
Columns: []clause.Column{
{Name: "used_amount"},
},
}).
Where("id = ?", id).
Updates(map[string]interface{}{
"used_amount": gorm.Expr("used_amount + ?", amount),
"request_count": gorm.Expr("request_count + ?", count),
})
return HandleUpdateResult(result, ErrGroupNotFound)
}
func UpdateGroupRPMRatio(id string, rpmRatio float64) (err error) {
defer func() {
if err == nil {
if err := CacheUpdateGroupRPMRatio(id, rpmRatio); err != nil {
log.Error("cache update group rpm failed: " + err.Error())
}
}
}()
result := DB.Model(&Group{}).Where("id = ?", id).Update("rpm_ratio", rpmRatio)
return HandleUpdateResult(result, ErrGroupNotFound)
}
func UpdateGroupRPM(id string, rpm map[string]int64) (err error) {
defer func() {
if err == nil {
if err := CacheUpdateGroupRPM(id, rpm); err != nil {
log.Error("cache update group rpm failed: " + err.Error())
}
}
}()
jsonRpm, err := json.Marshal(rpm)
if err != nil {
return err
}
result := DB.Model(&Group{}).Where("id = ?", id).Update("rpm", jsonRpm)
return HandleUpdateResult(result, ErrGroupNotFound)
}
func UpdateGroupTPMRatio(id string, tpmRatio float64) (err error) {
defer func() {
if err == nil {
if err := CacheUpdateGroupTPMRatio(id, tpmRatio); err != nil {
log.Error("cache update group tpm ratio failed: " + err.Error())
}
}
}()
result := DB.Model(&Group{}).Where("id = ?", id).Update("tpm_ratio", tpmRatio)
return HandleUpdateResult(result, ErrGroupNotFound)
}
func UpdateGroupTPM(id string, tpm map[string]int64) (err error) {
defer func() {
if err == nil {
if err := CacheUpdateGroupTPM(id, tpm); err != nil {
log.Error("cache update group tpm failed: " + err.Error())
}
}
}()
jsonTpm, err := json.Marshal(tpm)
if err != nil {
return err
}
result := DB.Model(&Group{}).Where("id = ?", id).Update("tpm", jsonTpm)
return HandleUpdateResult(result, ErrGroupNotFound)
}
@@ -202,19 +233,6 @@ func SearchGroup(keyword string, startIdx int, num int, order string, status int
} else {
tx = tx.Where("id LIKE ?", "%"+keyword+"%")
}
if keyword != "" {
var conditions []string
var values []interface{}
if status == 0 {
conditions = append(conditions, "status = ?")
values = append(values, 1)
}
if len(conditions) > 0 {
tx = tx.Where(fmt.Sprintf("(%s)", strings.Join(conditions, " OR ")), values...)
}
}
err = tx.Count(&total).Error
if err != nil {
return nil, 0, err
File diff suppressed because it is too large Load Diff
+12 -4
View File
@@ -3,6 +3,7 @@ package model
import (
"fmt"
"os"
"path/filepath"
"strings"
"time"
@@ -88,8 +89,15 @@ func openMySQL(dsn string) (*gorm.DB, error) {
}
func openSQLite() (*gorm.DB, error) {
log.Info("SQL_DSN not set, using SQLite as database")
log.Info("SQL_DSN not set, using SQLite as database: ", common.SQLitePath)
common.UsingSQLite = true
baseDir := filepath.Dir(common.SQLitePath)
if err := os.MkdirAll(baseDir, 0o755); err != nil {
log.Fatal("failed to create base directory: " + err.Error())
return nil, err
}
dsn := fmt.Sprintf("%s?_busy_timeout=%d", common.SQLitePath, common.SQLiteBusyTimeout)
return gorm.Open(sqlite.Open(dsn), &gorm.Config{
PrepareStmt: true, // precompile SQL
@@ -194,9 +202,9 @@ func setDBConns(db *gorm.DB) {
return
}
sqlDB.SetMaxIdleConns(env.Int("SQL_MAX_IDLE_CONNS", 100))
sqlDB.SetMaxOpenConns(env.Int("SQL_MAX_OPEN_CONNS", 1000))
sqlDB.SetConnMaxLifetime(time.Second * time.Duration(env.Int("SQL_MAX_LIFETIME", 60)))
sqlDB.SetMaxIdleConns(int(env.Int64("SQL_MAX_IDLE_CONNS", 100)))
sqlDB.SetMaxOpenConns(int(env.Int64("SQL_MAX_OPEN_CONNS", 1000)))
sqlDB.SetConnMaxLifetime(time.Second * time.Duration(env.Int64("SQL_MAX_LIFETIME", 60)))
}
func closeDB(db *gorm.DB) error {
+43 -51
View File
@@ -10,52 +10,9 @@ import (
"gorm.io/gorm"
)
//nolint:revive
type ModelConfigKey string
const (
ModelConfigMaxContextTokensKey ModelConfigKey = "max_context_tokens"
ModelConfigMaxInputTokensKey ModelConfigKey = "max_input_tokens"
ModelConfigMaxOutputTokensKey ModelConfigKey = "max_output_tokens"
ModelConfigToolChoiceKey ModelConfigKey = "tool_choice"
ModelConfigFunctionCallingKey ModelConfigKey = "function_calling"
ModelConfigSupportFormatsKey ModelConfigKey = "support_formats"
ModelConfigSupportVoicesKey ModelConfigKey = "support_voices"
)
//nolint:revive
type ModelOwner string
const (
ModelOwnerOpenAI ModelOwner = "openai"
ModelOwnerAlibaba ModelOwner = "alibaba"
ModelOwnerTencent ModelOwner = "tencent"
ModelOwnerXunfei ModelOwner = "xunfei"
ModelOwnerDeepSeek ModelOwner = "deepseek"
ModelOwnerMoonshot ModelOwner = "moonshot"
ModelOwnerMiniMax ModelOwner = "minimax"
ModelOwnerBaidu ModelOwner = "baidu"
ModelOwnerGoogle ModelOwner = "google"
ModelOwnerBAAI ModelOwner = "baai"
ModelOwnerFunAudioLLM ModelOwner = "funaudiollm"
ModelOwnerDoubao ModelOwner = "doubao"
ModelOwnerFishAudio ModelOwner = "fishaudio"
ModelOwnerChatGLM ModelOwner = "chatglm"
ModelOwnerStabilityAI ModelOwner = "stabilityai"
ModelOwnerNetease ModelOwner = "netease"
ModelOwnerAI360 ModelOwner = "ai360"
ModelOwnerAnthropic ModelOwner = "anthropic"
ModelOwnerMeta ModelOwner = "meta"
ModelOwnerBaichuan ModelOwner = "baichuan"
ModelOwnerMistral ModelOwner = "mistral"
ModelOwnerOpenChat ModelOwner = "openchat"
ModelOwnerMicrosoft ModelOwner = "microsoft"
ModelOwnerDefog ModelOwner = "defog"
ModelOwnerNexusFlow ModelOwner = "nexusflow"
ModelOwnerCohere ModelOwner = "cohere"
ModelOwnerHuggingFace ModelOwner = "huggingface"
ModelOwnerLingyiWanwu ModelOwner = "lingyiwanwu"
ModelOwnerStepFun ModelOwner = "stepfun"
// /1K tokens
PriceUnit = 1000
)
//nolint:revive
@@ -63,14 +20,21 @@ type ModelConfig struct {
CreatedAt time.Time `gorm:"index;autoCreateTime" json:"created_at"`
UpdatedAt time.Time `gorm:"index;autoUpdateTime" json:"updated_at"`
Config map[ModelConfigKey]any `gorm:"serializer:fastjson;type:text" json:"config,omitempty"`
ImagePrices map[string]float64 `gorm:"serializer:fastjson" json:"image_prices"`
ImagePrices map[string]float64 `gorm:"serializer:fastjson" json:"image_prices,omitempty"`
Model string `gorm:"primaryKey" json:"model"`
Owner ModelOwner `gorm:"type:varchar(255);index" json:"owner"`
ImageMaxBatchSize int `json:"image_batch_size"`
// relaymode/define.go
Type int `json:"type"`
InputPrice float64 `json:"input_price"`
OutputPrice float64 `json:"output_price"`
ImageMaxBatchSize int `json:"image_batch_size,omitempty"`
Type int `json:"type"` // relaymode/define.go
InputPrice float64 `json:"input_price,omitempty"`
OutputPrice float64 `json:"output_price,omitempty"`
RPM int64 `json:"rpm,omitempty"`
TPM int64 `json:"tpm,omitempty"`
}
func NewDefaultModelConfig(model string) *ModelConfig {
return &ModelConfig{
Model: model,
}
}
func (c *ModelConfig) MarshalJSON() ([]byte, error) {
@@ -86,6 +50,34 @@ func (c *ModelConfig) MarshalJSON() ([]byte, error) {
})
}
func (c *ModelConfig) MaxContextTokens() (int, bool) {
return GetModelConfigInt(c.Config, ModelConfigMaxContextTokensKey)
}
func (c *ModelConfig) MaxInputTokens() (int, bool) {
return GetModelConfigInt(c.Config, ModelConfigMaxInputTokensKey)
}
func (c *ModelConfig) MaxOutputTokens() (int, bool) {
return GetModelConfigInt(c.Config, ModelConfigMaxOutputTokensKey)
}
func (c *ModelConfig) SupportVision() (bool, bool) {
return GetModelConfigBool(c.Config, ModelConfigVisionKey)
}
func (c *ModelConfig) SupportVoices() ([]string, bool) {
return GetModelConfigStringSlice(c.Config, ModelConfigSupportVoicesKey)
}
func (c *ModelConfig) SupportToolChoice() (bool, bool) {
return GetModelConfigBool(c.Config, ModelConfigToolChoiceKey)
}
func (c *ModelConfig) SupportFormats() ([]string, bool) {
return GetModelConfigStringSlice(c.Config, ModelConfigSupportFormatsKey)
}
func GetModelConfigs(startIdx int, num int, model string) (configs []*ModelConfig, total int64, err error) {
tx := DB.Model(&ModelConfig{})
if model != "" {
+129 -57
View File
@@ -11,7 +11,6 @@ import (
"time"
json "github.com/json-iterator/go"
"github.com/labring/sealos/service/aiproxy/common/config"
"github.com/labring/sealos/service/aiproxy/common/conv"
log "github.com/sirupsen/logrus"
@@ -24,41 +23,78 @@ type Option struct {
func GetAllOption() ([]*Option, error) {
var options []*Option
err := DB.Find(&options).Error
err := DB.Where("key IN (?)", optionKeys).Find(&options).Error
return options, err
}
func InitOptionMap() error {
config.OptionMapRWMutex.Lock()
config.OptionMap = make(map[string]string)
config.OptionMap["LogDetailStorageHours"] = strconv.FormatInt(config.GetLogDetailStorageHours(), 10)
config.OptionMap["DisableServe"] = strconv.FormatBool(config.GetDisableServe())
config.OptionMap["AutomaticDisableChannelEnabled"] = strconv.FormatBool(config.GetAutomaticDisableChannelEnabled())
config.OptionMap["AutomaticEnableChannelWhenTestSucceedEnabled"] = strconv.FormatBool(config.GetAutomaticEnableChannelWhenTestSucceedEnabled())
config.OptionMap["ApproximateTokenEnabled"] = strconv.FormatBool(config.GetApproximateTokenEnabled())
config.OptionMap["BillingEnabled"] = strconv.FormatBool(config.GetBillingEnabled())
config.OptionMap["RetryTimes"] = strconv.FormatInt(config.GetRetryTimes(), 10)
config.OptionMap["GlobalApiRateLimitNum"] = strconv.FormatInt(config.GetGlobalAPIRateLimitNum(), 10)
config.OptionMap["DefaultGroupQPM"] = strconv.FormatInt(config.GetDefaultGroupQPM(), 10)
defaultChannelModelsJSON, _ := json.Marshal(config.GetDefaultChannelModels())
config.OptionMap["DefaultChannelModels"] = conv.BytesToString(defaultChannelModelsJSON)
defaultChannelModelMappingJSON, _ := json.Marshal(config.GetDefaultChannelModelMapping())
config.OptionMap["DefaultChannelModelMapping"] = conv.BytesToString(defaultChannelModelMappingJSON)
config.OptionMap["GeminiSafetySetting"] = config.GetGeminiSafetySetting()
config.OptionMap["GeminiVersion"] = config.GetGeminiVersion()
config.OptionMap["GroupMaxTokenNum"] = strconv.FormatInt(int64(config.GetGroupMaxTokenNum()), 10)
config.OptionMapRWMutex.Unlock()
err := loadOptionsFromDatabase(true)
func GetOption(key string) (*Option, error) {
if !slices.Contains(optionKeys, key) {
return nil, ErrUnknownOptionKey
}
var option Option
err := DB.Where("key = ?", key).First(&option).Error
return &option, err
}
var (
optionMap = make(map[string]string)
// allowed option keys
optionKeys []string
)
func InitOption2DB() error {
err := initOptionMap()
if err != nil {
return err
}
err = loadOptionsFromDatabase(true)
if err != nil {
return err
}
return storeOptionMap()
}
func initOptionMap() error {
optionMap["LogDetailStorageHours"] = strconv.FormatInt(config.GetLogDetailStorageHours(), 10)
optionMap["DisableServe"] = strconv.FormatBool(config.GetDisableServe())
optionMap["BillingEnabled"] = strconv.FormatBool(config.GetBillingEnabled())
optionMap["RetryTimes"] = strconv.FormatInt(config.GetRetryTimes(), 10)
optionMap["ModelErrorAutoBanRate"] = strconv.FormatFloat(config.GetModelErrorAutoBanRate(), 'f', -1, 64)
optionMap["EnableModelErrorAutoBan"] = strconv.FormatBool(config.GetEnableModelErrorAutoBan())
timeoutWithModelTypeJSON, err := json.Marshal(config.GetTimeoutWithModelType())
if err != nil {
return err
}
optionMap["TimeoutWithModelType"] = conv.BytesToString(timeoutWithModelTypeJSON)
defaultChannelModelsJSON, err := json.Marshal(config.GetDefaultChannelModels())
if err != nil {
return err
}
optionMap["DefaultChannelModels"] = conv.BytesToString(defaultChannelModelsJSON)
defaultChannelModelMappingJSON, err := json.Marshal(config.GetDefaultChannelModelMapping())
if err != nil {
return err
}
optionMap["DefaultChannelModelMapping"] = conv.BytesToString(defaultChannelModelMappingJSON)
optionMap["GeminiSafetySetting"] = config.GetGeminiSafetySetting()
optionMap["GroupMaxTokenNum"] = strconv.FormatInt(config.GetGroupMaxTokenNum(), 10)
groupConsumeLevelRatioJSON, err := json.Marshal(config.GetGroupConsumeLevelRatio())
if err != nil {
return err
}
optionMap["GroupConsumeLevelRatio"] = conv.BytesToString(groupConsumeLevelRatioJSON)
optionMap["InternalToken"] = config.GetInternalToken()
optionKeys = make([]string, 0, len(optionMap))
for key := range optionMap {
optionKeys = append(optionKeys, key)
}
return nil
}
func storeOptionMap() error {
config.OptionMapRWMutex.Lock()
defer config.OptionMapRWMutex.Unlock()
for key, value := range config.OptionMap {
for key, value := range optionMap {
err := saveOption(key, value)
if err != nil {
return err
@@ -73,9 +109,18 @@ func loadOptionsFromDatabase(isInit bool) error {
return err
}
for _, option := range options {
err := updateOptionMap(option.Key, option.Value, isInit)
if err != nil && !errors.Is(err, ErrUnknownOptionKey) {
log.Errorf("failed to update option: %s, value: %s, error: %s", option.Key, option.Value, err.Error())
err := updateOption(option.Key, option.Value, isInit)
if err != nil {
if !errors.Is(err, ErrUnknownOptionKey) {
return fmt.Errorf("failed to update option: %s, value: %s, error: %w", option.Key, option.Value, err)
}
if isInit {
log.Warnf("unknown option: %s, value: %s", option.Key, option.Value)
}
continue
}
if isInit {
delete(optionMap, option.Key)
}
}
return nil
@@ -91,7 +136,7 @@ func SyncOptions(ctx context.Context, wg *sync.WaitGroup, frequency time.Duratio
case <-ctx.Done():
return
case <-ticker.C:
if err := loadOptionsFromDatabase(true); err != nil {
if err := loadOptionsFromDatabase(false); err != nil {
log.Error("failed to sync options from database: " + err.Error())
}
}
@@ -108,7 +153,7 @@ func saveOption(key string, value string) error {
}
func UpdateOption(key string, value string) error {
err := updateOptionMap(key, value, false)
err := updateOption(key, value, false)
if err != nil {
return err
}
@@ -136,25 +181,22 @@ func isTrue(value string) bool {
return result
}
func updateOptionMap(key string, value string, isInit bool) (err error) {
config.OptionMapRWMutex.Lock()
defer config.OptionMapRWMutex.Unlock()
config.OptionMap[key] = value
//nolint:gocyclo
func updateOption(key string, value string, isInit bool) (err error) {
switch key {
case "InternalToken":
config.SetInternalToken(value)
case "LogDetailStorageHours":
logDetailStorageHours, err := strconv.ParseInt(value, 10, 64)
if err != nil {
return err
}
if logDetailStorageHours < 0 {
return errors.New("log detail storage hours must be greater than 0")
}
config.SetLogDetailStorageHours(logDetailStorageHours)
case "DisableServe":
config.SetDisableServe(isTrue(value))
case "AutomaticDisableChannelEnabled":
config.SetAutomaticDisableChannelEnabled(isTrue(value))
case "AutomaticEnableChannelWhenTestSucceedEnabled":
config.SetAutomaticEnableChannelWhenTestSucceedEnabled(isTrue(value))
case "ApproximateTokenEnabled":
config.SetApproximateTokenEnabled(isTrue(value))
case "BillingEnabled":
config.SetBillingEnabled(isTrue(value))
case "GroupMaxTokenNum":
@@ -162,23 +204,12 @@ func updateOptionMap(key string, value string, isInit bool) (err error) {
if err != nil {
return err
}
config.SetGroupMaxTokenNum(int32(groupMaxTokenNum))
if groupMaxTokenNum < 0 {
return errors.New("group max token num must be greater than 0")
}
config.SetGroupMaxTokenNum(groupMaxTokenNum)
case "GeminiSafetySetting":
config.SetGeminiSafetySetting(value)
case "GeminiVersion":
config.SetGeminiVersion(value)
case "GlobalApiRateLimitNum":
globalAPIRateLimitNum, err := strconv.ParseInt(value, 10, 64)
if err != nil {
return err
}
config.SetGlobalAPIRateLimitNum(globalAPIRateLimitNum)
case "DefaultGroupQPM":
defaultGroupQPM, err := strconv.ParseInt(value, 10, 64)
if err != nil {
return err
}
config.SetDefaultGroupQPM(defaultGroupQPM)
case "DefaultChannelModels":
var newModels map[int][]string
err := json.Unmarshal(conv.StringToBytes(value), &newModels)
@@ -196,7 +227,7 @@ func updateOptionMap(key string, value string, isInit bool) (err error) {
for model := range allModelsMap {
allModels = append(allModels, model)
}
foundModels, missingModels, err := CheckModelConfig(allModels)
foundModels, missingModels, err := GetModelConfigWithModels(allModels)
if err != nil {
return err
}
@@ -229,7 +260,48 @@ func updateOptionMap(key string, value string, isInit bool) (err error) {
if err != nil {
return err
}
if retryTimes < 0 {
return errors.New("retry times must be greater than 0")
}
config.SetRetryTimes(retryTimes)
case "EnableModelErrorAutoBan":
config.SetEnableModelErrorAutoBan(isTrue(value))
case "ModelErrorAutoBanRate":
modelErrorAutoBanRate, err := strconv.ParseFloat(value, 64)
if err != nil {
return err
}
if modelErrorAutoBanRate < 0 || modelErrorAutoBanRate > 1 {
return errors.New("model error auto ban rate must be between 0 and 1")
}
config.SetModelErrorAutoBanRate(modelErrorAutoBanRate)
case "TimeoutWithModelType":
var newTimeoutWithModelType map[int]int64
err := json.Unmarshal(conv.StringToBytes(value), &newTimeoutWithModelType)
if err != nil {
return err
}
for _, v := range newTimeoutWithModelType {
if v < 0 {
return errors.New("timeout must be greater than 0")
}
}
config.SetTimeoutWithModelType(newTimeoutWithModelType)
case "GroupConsumeLevelRatio":
var newGroupRpmRatio map[float64]float64
err := json.Unmarshal(conv.StringToBytes(value), &newGroupRpmRatio)
if err != nil {
return err
}
for k, v := range newGroupRpmRatio {
if k < 0 {
return errors.New("consume level must be greater than 0")
}
if v < 0 {
return errors.New("rpm ratio must be greater than 0")
}
}
config.SetGroupConsumeLevelRatio(newGroupRpmRatio)
default:
return ErrUnknownOptionKey
}
+36
View File
@@ -0,0 +1,36 @@
package model
//nolint:revive
type ModelOwner string
const (
ModelOwnerOpenAI ModelOwner = "openai"
ModelOwnerAlibaba ModelOwner = "alibaba"
ModelOwnerTencent ModelOwner = "tencent"
ModelOwnerXunfei ModelOwner = "xunfei"
ModelOwnerDeepSeek ModelOwner = "deepseek"
ModelOwnerMoonshot ModelOwner = "moonshot"
ModelOwnerMiniMax ModelOwner = "minimax"
ModelOwnerBaidu ModelOwner = "baidu"
ModelOwnerGoogle ModelOwner = "google"
ModelOwnerBAAI ModelOwner = "baai"
ModelOwnerFunAudioLLM ModelOwner = "funaudiollm"
ModelOwnerDoubao ModelOwner = "doubao"
ModelOwnerFishAudio ModelOwner = "fishaudio"
ModelOwnerChatGLM ModelOwner = "chatglm"
ModelOwnerStabilityAI ModelOwner = "stabilityai"
ModelOwnerNetease ModelOwner = "netease"
ModelOwnerAI360 ModelOwner = "ai360"
ModelOwnerAnthropic ModelOwner = "anthropic"
ModelOwnerMeta ModelOwner = "meta"
ModelOwnerBaichuan ModelOwner = "baichuan"
ModelOwnerMistral ModelOwner = "mistral"
ModelOwnerOpenChat ModelOwner = "openchat"
ModelOwnerMicrosoft ModelOwner = "microsoft"
ModelOwnerDefog ModelOwner = "defog"
ModelOwnerNexusFlow ModelOwner = "nexusflow"
ModelOwnerCohere ModelOwner = "cohere"
ModelOwnerHuggingFace ModelOwner = "huggingface"
ModelOwnerLingyiWanwu ModelOwner = "lingyiwanwu"
ModelOwnerStepFun ModelOwner = "stepfun"
)
+33 -181
View File
@@ -6,8 +6,6 @@ import (
"strings"
"time"
json "github.com/json-iterator/go"
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/common/config"
log "github.com/sirupsen/logrus"
@@ -29,7 +27,6 @@ const (
type Token struct {
CreatedAt time.Time `json:"created_at"`
ExpiredAt time.Time `json:"expired_at"`
AccessedAt time.Time `gorm:"index" json:"accessed_at"`
Group *Group `gorm:"foreignKey:GroupID" json:"-"`
Key string `gorm:"type:char(48);uniqueIndex" json:"key"`
Name EmptyNullString `gorm:"index;uniqueIndex:idx_group_name;not null" json:"name"`
@@ -43,26 +40,11 @@ type Token struct {
RequestCount int `gorm:"index" json:"request_count"`
}
func (t *Token) MarshalJSON() ([]byte, error) {
type Alias Token
return json.Marshal(&struct {
*Alias
CreatedAt int64 `json:"created_at"`
AccessedAt int64 `json:"accessed_at"`
ExpiredAt int64 `json:"expired_at"`
}{
Alias: (*Alias)(t),
CreatedAt: t.CreatedAt.UnixMilli(),
AccessedAt: t.AccessedAt.UnixMilli(),
ExpiredAt: t.ExpiredAt.UnixMilli(),
})
}
//nolint:goconst
func getTokenOrder(order string) string {
prefix, suffix, _ := strings.Cut(order, "-")
switch prefix {
case "name", "accessed_at", "expired_at", "group", "used_amount", "request_count", "id", "created_at":
case "name", "expired_at", "group", "used_amount", "request_count", "id", "created_at":
switch suffix {
case "asc":
return prefix + " asc"
@@ -91,7 +73,7 @@ func InsertToken(token *Token, autoCreateGroup bool) error {
if err != nil {
return err
}
if count >= int64(maxTokenNum) {
if count >= maxTokenNum {
return errors.New("group max token num reached")
}
}
@@ -106,34 +88,11 @@ func InsertToken(token *Token, autoCreateGroup bool) error {
return nil
}
func GetTokens(startIdx int, num int, order string, group string, status int) (tokens []*Token, total int64, err error) {
func GetTokens(group string, startIdx int, num int, order string, status int) (tokens []*Token, total int64, err error) {
tx := DB.Model(&Token{})
if group != "" {
tx = tx.Where("group_id = ?", group)
}
if status != 0 {
tx = tx.Where("status = ?", status)
}
err = tx.Count(&total).Error
if err != nil {
return nil, 0, err
}
if total <= 0 {
return nil, 0, nil
}
err = tx.Order(getTokenOrder(order)).Limit(num).Offset(startIdx).Find(&tokens).Error
return tokens, total, err
}
func GetGroupTokens(group string, startIdx int, num int, order string, status int) (tokens []*Token, total int64, err error) {
if group == "" {
return nil, 0, errors.New("group is empty")
}
tx := DB.Model(&Token{}).Where("group_id = ?", group)
if status != 0 {
tx = tx.Where("status = ?", status)
@@ -151,7 +110,7 @@ func GetGroupTokens(group string, startIdx int, num int, order string, status in
return tokens, total, err
}
func SearchTokens(keyword string, startIdx int, num int, order string, status int, name string, key string, group string) (tokens []*Token, total int64, err error) {
func SearchTokens(group string, keyword string, startIdx int, num int, order string, status int, name string, key string) (tokens []*Token, total int64, err error) {
tx := DB.Model(&Token{})
if group != "" {
tx = tx.Where("group_id = ?", group)
@@ -169,88 +128,31 @@ func SearchTokens(keyword string, startIdx int, num int, order string, status in
if keyword != "" {
var conditions []string
var values []interface{}
if status == 0 {
conditions = append(conditions, "status = ?")
values = append(values, 1)
}
if group == "" {
if common.UsingPostgreSQL {
conditions = append(conditions, "group_id ILIKE ?")
} else {
conditions = append(conditions, "group_id LIKE ?")
}
values = append(values, "%"+keyword+"%")
}
if name == "" {
if common.UsingPostgreSQL {
conditions = append(conditions, "name ILIKE ?")
} else {
conditions = append(conditions, "name LIKE ?")
}
values = append(values, "%"+keyword+"%")
}
if key == "" {
if common.UsingPostgreSQL {
conditions = append(conditions, "key ILIKE ?")
} else {
conditions = append(conditions, "key LIKE ?")
}
conditions = append(conditions, "group_id = ?")
values = append(values, keyword)
}
if len(conditions) > 0 {
tx = tx.Where(fmt.Sprintf("(%s)", strings.Join(conditions, " OR ")), values...)
}
}
err = tx.Count(&total).Error
if err != nil {
return nil, 0, err
}
if total <= 0 {
return nil, 0, nil
}
err = tx.Order(getTokenOrder(order)).Limit(num).Offset(startIdx).Find(&tokens).Error
return tokens, total, err
}
func SearchGroupTokens(group string, keyword string, startIdx int, num int, order string, status int, name string, key string) (tokens []*Token, total int64, err error) {
if group == "" {
return nil, 0, errors.New("group is empty")
}
tx := DB.Model(&Token{}).Where("group_id = ?", group)
if status != 0 {
tx = tx.Where("status = ?", status)
}
if name != "" {
tx = tx.Where("name = ?", name)
}
if key != "" {
tx = tx.Where("key = ?", key)
}
if keyword != "" {
var conditions []string
var values []interface{}
if status == 0 {
conditions = append(conditions, "status = ?")
values = append(values, 1)
values = append(values, String2Int(keyword))
}
if name == "" {
if common.UsingPostgreSQL {
conditions = append(conditions, "name ILIKE ?")
} else {
conditions = append(conditions, "name LIKE ?")
}
values = append(values, "%"+keyword+"%")
}
if key == "" {
if common.UsingPostgreSQL {
conditions = append(conditions, "key ILIKE ?")
} else {
conditions = append(conditions, "key LIKE ?")
}
conditions = append(conditions, "name = ?")
values = append(values, keyword)
}
if key == "" {
conditions = append(conditions, "key = ?")
values = append(values, keyword)
}
if common.UsingPostgreSQL {
conditions = append(conditions, "models ILIKE ?")
} else {
conditions = append(conditions, "models LIKE ?")
}
values = append(values, "%"+keyword+"%")
if len(conditions) > 0 {
tx = tx.Where(fmt.Sprintf("(%s)", strings.Join(conditions, " OR ")), values...)
}
@@ -291,10 +193,10 @@ func ValidateAndGetToken(key string) (token *TokenCache, err error) {
}
token, err = CacheGetTokenByKey(key)
if err != nil {
log.Error("get token from cache failed: " + err.Error())
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New("invalid token")
}
log.Error("get token from cache failed: " + err.Error())
return nil, errors.New("token validation failed")
}
switch token.Status {
@@ -302,12 +204,14 @@ func ValidateAndGetToken(key string) (token *TokenCache, err error) {
return nil, fmt.Errorf("token (%s[%d]) quota is exhausted", token.Name, token.ID)
case TokenStatusExpired:
return nil, fmt.Errorf("token (%s[%d]) is expired", token.Name, token.ID)
case TokenStatusDisabled:
return nil, fmt.Errorf("token (%s[%d]) is disabled", token.Name, token.ID)
}
if token.Status != TokenStatusEnabled {
return nil, fmt.Errorf("token (%s[%d]) is not available", token.Name, token.ID)
}
if !time.Time(token.ExpiredAt).IsZero() && time.Time(token.ExpiredAt).Before(time.Now()) {
err := UpdateTokenStatusAndAccessedAt(token.ID, TokenStatusExpired)
err := UpdateTokenStatus(token.ID, TokenStatusExpired)
if err != nil {
log.Error("failed to update token status" + err.Error())
}
@@ -315,7 +219,7 @@ func ValidateAndGetToken(key string) (token *TokenCache, err error) {
}
if token.Quota > 0 && token.UsedAmount >= token.Quota {
// in this case, we can make sure the token is exhausted
err := UpdateTokenStatusAndAccessedAt(token.ID, TokenStatusExhausted)
err := UpdateTokenStatus(token.ID, TokenStatusExhausted)
if err != nil {
log.Error("failed to update token status" + err.Error())
}
@@ -348,8 +252,8 @@ func UpdateTokenStatus(id int, status int) (err error) {
token := Token{ID: id}
defer func() {
if err == nil {
if err := CacheDeleteToken(token.Key); err != nil {
log.Error("delete token from cache failed: " + err.Error())
if err := CacheUpdateTokenStatus(token.Key, status); err != nil {
log.Error("update token status in cache failed: " + err.Error())
}
}
}()
@@ -369,63 +273,12 @@ func UpdateTokenStatus(id int, status int) (err error) {
return HandleUpdateResult(result, ErrTokenNotFound)
}
func UpdateTokenStatusAndAccessedAt(id int, status int) (err error) {
token := Token{ID: id}
defer func() {
if err == nil {
if err := CacheDeleteToken(token.Key); err != nil {
log.Error("delete token from cache failed: " + err.Error())
}
}
}()
result := DB.
Model(&token).
Clauses(clause.Returning{
Columns: []clause.Column{
{Name: "key"},
},
}).
Where("id = ?", id).Updates(
map[string]interface{}{
"status": status,
"accessed_at": time.Now(),
},
)
return HandleUpdateResult(result, ErrTokenNotFound)
}
func UpdateGroupTokenStatusAndAccessedAt(group string, id int, status int) (err error) {
token := Token{}
defer func() {
if err == nil {
if err := CacheDeleteToken(token.Key); err != nil {
log.Error("delete token from cache failed: " + err.Error())
}
}
}()
result := DB.
Model(&token).
Clauses(clause.Returning{
Columns: []clause.Column{
{Name: "key"},
},
}).
Where("id = ? and group_id = ?", id, group).
Updates(
map[string]interface{}{
"status": status,
"accessed_at": time.Now(),
},
)
return HandleUpdateResult(result, ErrTokenNotFound)
}
func UpdateGroupTokenStatus(group string, id int, status int) (err error) {
token := Token{}
defer func() {
if err == nil {
if err := CacheDeleteToken(token.Key); err != nil {
log.Error("delete token from cache failed: " + err.Error())
if err := CacheUpdateTokenStatus(token.Key, status); err != nil {
log.Error("update token status in cache failed: " + err.Error())
}
}
}()
@@ -586,7 +439,6 @@ func UpdateTokenUsedAmount(id int, amount float64, requestCount int) (err error)
map[string]interface{}{
"used_amount": gorm.Expr("used_amount + ?", amount),
"request_count": gorm.Expr("request_count + ?", requestCount),
"accessed_at": time.Now(),
},
)
return HandleUpdateResult(result, ErrTokenNotFound)
@@ -596,8 +448,8 @@ func UpdateTokenName(id int, name string) (err error) {
token := &Token{ID: id}
defer func() {
if err == nil {
if err := CacheDeleteToken(token.Key); err != nil {
log.Error("delete token from cache failed: " + err.Error())
if err := CacheUpdateTokenName(token.Key, name); err != nil {
log.Error("update token name in cache failed: " + err.Error())
}
}
}()
@@ -620,8 +472,8 @@ func UpdateGroupTokenName(group string, id int, name string) (err error) {
token := &Token{ID: id, GroupID: group}
defer func() {
if err == nil {
if err := CacheDeleteToken(token.Key); err != nil {
log.Error("delete token from cache failed: " + err.Error())
if err := CacheUpdateTokenName(token.Key, name); err != nil {
log.Error("update token name in cache failed: " + err.Error())
}
}
}()
+29 -9
View File
@@ -4,6 +4,7 @@ import (
"database/sql/driver"
"errors"
"fmt"
"strconv"
"strings"
"time"
@@ -58,6 +59,7 @@ func BatchRecordConsume(
endpoint string,
content string,
mode int,
ip string,
requestDetail *RequestDetail,
) error {
errs := []error{}
@@ -78,22 +80,29 @@ func BatchRecordConsume(
endpoint,
content,
mode,
ip,
requestDetail,
)
if err != nil {
errs = append(errs, fmt.Errorf("failed to record log: %w", err))
}
err = UpdateGroupUsedAmountAndRequestCount(group, amount, 1)
if err != nil {
errs = append(errs, fmt.Errorf("failed to update group used amount and request count: %w", err))
if group != "" {
err = UpdateGroupUsedAmountAndRequestCount(group, amount, 1)
if err != nil {
errs = append(errs, fmt.Errorf("failed to update group used amount and request count: %w", err))
}
}
err = UpdateTokenUsedAmount(tokenID, amount, 1)
if err != nil {
errs = append(errs, fmt.Errorf("failed to update token used amount: %w", err))
if tokenID > 0 {
err = UpdateTokenUsedAmount(tokenID, amount, 1)
if err != nil {
errs = append(errs, fmt.Errorf("failed to update token used amount: %w", err))
}
}
err = UpdateChannelUsedAmount(channelID, amount, 1)
if err != nil {
errs = append(errs, fmt.Errorf("failed to update channel used amount: %w", err))
if channelID > 0 {
err = UpdateChannelUsedAmount(channelID, amount, 1)
if err != nil {
errs = append(errs, fmt.Errorf("failed to update channel used amount: %w", err))
}
}
if len(errs) == 0 {
return nil
@@ -131,3 +140,14 @@ func (ns EmptyNullString) Value() (driver.Value, error) {
}
return string(ns), nil
}
func String2Int(keyword string) int {
if keyword == "" {
return 0
}
i, err := strconv.Atoi(keyword)
if err != nil {
return 0
}
return i
}
+386
View File
@@ -0,0 +1,386 @@
package monitor
import (
"context"
"fmt"
"strconv"
"strings"
"time"
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/common/config"
"github.com/redis/go-redis/v9"
log "github.com/sirupsen/logrus"
)
// Redis key prefixes and patterns
const (
modelKeyPrefix = "model:"
bannedKeySuffix = ":banned"
statsKeySuffix = ":stats"
channelKeyPart = ":channel:"
)
// Redis scripts
var (
addRequestScript = redis.NewScript(addRequestLuaScript)
getChannelModelErrorRateScript = redis.NewScript(getChannelModelErrorRateLuaScript)
getBannedChannelsScript = redis.NewScript(getBannedChannelsLuaScript)
clearChannelModelErrorsScript = redis.NewScript(clearChannelModelErrorsLuaScript)
clearChannelAllModelErrorsScript = redis.NewScript(clearChannelAllModelErrorsLuaScript)
clearAllModelErrorsScript = redis.NewScript(clearAllModelErrorsLuaScript)
)
// Helper functions
func isFeatureEnabled() bool {
return common.RedisEnabled && config.GetEnableModelErrorAutoBan()
}
func buildStatsKey(model string, channelID interface{}) string {
return fmt.Sprintf("%s%s%s%v%s", modelKeyPrefix, model, channelKeyPart, channelID, statsKeySuffix)
}
// AddRequest adds a request record and checks if channel should be banned
func AddRequest(ctx context.Context, model string, channelID int64, isError bool) error {
if !isFeatureEnabled() {
return nil
}
errorFlag := 0
if isError {
errorFlag = 1
}
now := time.Now().UnixMilli()
val, err := addRequestScript.Run(
ctx,
common.RDB,
[]string{model},
channelID,
errorFlag,
now,
config.GetModelErrorAutoBanRate(),
time.Second.Milliseconds()*15,
).Int64()
if err != nil {
return err
}
log.Debugf("add request result: %d", val)
if val == 1 {
log.Errorf("channel %d model %s is banned", channelID, model)
}
return nil
}
// GetChannelModelErrorRates gets error rates for a specific channel
func GetChannelModelErrorRates(ctx context.Context, channelID int64) (map[string]float64, error) {
if !isFeatureEnabled() {
return nil, nil
}
result := make(map[string]float64)
pattern := buildStatsKey("*", channelID)
now := time.Now().UnixMilli()
iter := common.RDB.Scan(ctx, 0, pattern, 0).Iterator()
for iter.Next(ctx) {
key := iter.Val()
parts := strings.Split(key, ":")
if len(parts) != 5 || parts[4] != "stats" {
continue
}
model := parts[1]
rate, err := getChannelModelErrorRateScript.Run(
ctx,
common.RDB,
[]string{key},
now,
).Float64()
if err != nil {
return nil, err
}
result[model] = rate
}
if err := iter.Err(); err != nil {
return nil, err
}
return result, nil
}
// GetBannedChannels gets banned channels for a specific model
func GetBannedChannels(ctx context.Context, model string) ([]int64, error) {
if !isFeatureEnabled() {
return nil, nil
}
result, err := getBannedChannelsScript.Run(ctx, common.RDB, []string{model}).Int64Slice()
if err != nil {
return nil, err
}
return result, nil
}
// ClearChannelModelErrors clears errors for a specific channel and model
func ClearChannelModelErrors(ctx context.Context, model string, channelID int) error {
if !isFeatureEnabled() {
return nil
}
return clearChannelModelErrorsScript.Run(
ctx,
common.RDB,
[]string{model},
strconv.Itoa(channelID),
).Err()
}
// ClearChannelAllModelErrors clears all errors for a specific channel
func ClearChannelAllModelErrors(ctx context.Context, channelID int) error {
if !isFeatureEnabled() {
return nil
}
return clearChannelAllModelErrorsScript.Run(
ctx,
common.RDB,
[]string{},
strconv.Itoa(channelID),
).Err()
}
// ClearAllModelErrors clears all error records
func ClearAllModelErrors(ctx context.Context) error {
if !isFeatureEnabled() {
return nil
}
return clearAllModelErrorsScript.Run(ctx, common.RDB, []string{}).Err()
}
// GetAllBannedChannels gets all banned channels for all models
func GetAllBannedChannels(ctx context.Context) (map[string][]int64, error) {
if !isFeatureEnabled() {
return nil, nil
}
result := make(map[string][]int64)
iter := common.RDB.Scan(ctx, 0, modelKeyPrefix+"*"+bannedKeySuffix, 0).Iterator()
for iter.Next(ctx) {
key := iter.Val()
model := strings.Split(key, ":")[1]
channels, err := getBannedChannelsScript.Run(
ctx,
common.RDB,
[]string{model},
).Int64Slice()
if err != nil {
return nil, err
}
result[model] = channels
}
if err := iter.Err(); err != nil {
return nil, err
}
return result, nil
}
// GetAllChannelModelErrorRates gets error rates for all channels and models
func GetAllChannelModelErrorRates(ctx context.Context) (map[int64]map[string]float64, error) {
if !isFeatureEnabled() {
return nil, nil
}
result := make(map[int64]map[string]float64)
pattern := modelKeyPrefix + "*" + channelKeyPart + "*" + statsKeySuffix
now := time.Now().UnixMilli()
iter := common.RDB.Scan(ctx, 0, pattern, 0).Iterator()
for iter.Next(ctx) {
key := iter.Val()
parts := strings.Split(key, ":")
if len(parts) != 5 || parts[4] != "stats" {
continue
}
model := parts[1]
channelID, err := strconv.ParseInt(parts[3], 10, 64)
if err != nil {
continue
}
rate, err := getChannelModelErrorRateScript.Run(
ctx,
common.RDB,
[]string{key},
now,
).Float64()
if err != nil {
return nil, err
}
if _, exists := result[channelID]; !exists {
result[channelID] = make(map[string]float64)
}
result[channelID][model] = rate
}
if err := iter.Err(); err != nil {
return nil, err
}
return result, nil
}
// Lua scripts
const (
addRequestLuaScript = `
local model = KEYS[1]
local channel_id = ARGV[1]
local is_error = tonumber(ARGV[2])
local now_ts = tonumber(ARGV[3])
local max_error_rate = tonumber(ARGV[4])
local statsExpiry = tonumber(ARGV[5])
local banned_key = "model:" .. model .. ":banned"
local stats_key = "model:" .. model .. ":channel:" .. channel_id .. ":stats"
local maxSliceCount = 6
local current_slice = math.floor(now_ts / 1000)
if redis.call("SISMEMBER", banned_key, channel_id) == 1 then
return 2
end
local function parse_req_err(value)
if not value then return 0, 0 end
local r, e = value:match("^(%d+):(%d+)$")
return tonumber(r) or 0, tonumber(e) or 0
end
local function update_current_slice()
local req, err = parse_req_err(redis.call("HGET", stats_key, current_slice))
req = req + 1
err = err + (is_error == 1 and 1 or 0)
redis.call("HSET", stats_key, current_slice, req .. ":" .. err)
redis.call("PEXPIRE", stats_key, statsExpiry)
return req, err
end
local function calculate_error_rate()
local total_req, total_err = 0, 0
local min_valid_slice = current_slice - maxSliceCount
local all_slices = redis.call("HGETALL", stats_key)
local to_delete = {}
for i = 1, #all_slices, 2 do
local slice = tonumber(all_slices[i])
if slice < min_valid_slice then
table.insert(to_delete, all_slices[i])
else
local req, err = parse_req_err(all_slices[i+1])
total_req = total_req + req
total_err = total_err + err
end
end
if #to_delete > 0 then
redis.call("HDEL", stats_key, unpack(to_delete))
end
return total_req, total_err
end
update_current_slice()
if is_error == 0 then
return 0
end
local total_req, total_err = calculate_error_rate()
if total_req >= 10 and (total_err / total_req) >= max_error_rate then
redis.call("SADD", banned_key, channel_id)
redis.call("DEL", stats_key)
return 1
end
return 0
`
getChannelModelErrorRateLuaScript = `
local stats_key = KEYS[1]
local now_ts = tonumber(ARGV[1])
local maxSliceCount = 6
local current_slice = math.floor(now_ts / 1000)
local min_valid_slice = current_slice - maxSliceCount
local function parse_req_err(value)
if not value then return 0, 0 end
local r, e = value:match("^(%d+):(%d+)$")
return tonumber(r) or 0, tonumber(e) or 0
end
local total_req, total_err = 0, 0
local all_slices = redis.call("HGETALL", stats_key)
for i = 1, #all_slices, 2 do
local slice = tonumber(all_slices[i])
if slice >= min_valid_slice then
local req, err = parse_req_err(all_slices[i+1])
total_req = total_req + req
total_err = total_err + err
end
end
if total_req == 0 then return 0 end
return string.format("%.2f", total_err / total_req)
`
getBannedChannelsLuaScript = `
local model = KEYS[1]
return redis.call("SMEMBERS", "model:" .. model .. ":banned")
`
clearChannelModelErrorsLuaScript = `
local model = KEYS[1]
local channel_id = ARGV[1]
local stats_key = "model:" .. model .. ":channel:" .. channel_id .. ":stats"
local banned_key = "model:" .. model .. ":banned"
redis.call("DEL", stats_key)
redis.call("SREM", banned_key, channel_id)
return redis.status_reply("ok")
`
clearChannelAllModelErrorsLuaScript = `
local channel_id = ARGV[1]
local pattern = "model:*:channel:" .. channel_id .. ":stats"
local keys = redis.call("KEYS", pattern)
for _, key in ipairs(keys) do
redis.call("DEL", key)
local model = string.match(key, "model:(.*):channel:")
if model then
redis.call("SREM", "model:"..model..":banned", channel_id)
end
end
return redis.status_reply("ok")
`
clearAllModelErrorsLuaScript = `
local function del_keys(pattern)
local keys = redis.call("KEYS", pattern)
if #keys > 0 then redis.call("DEL", unpack(keys)) end
end
del_keys("model:*:channel:*:stats")
del_keys("model:*:banned")
return redis.status_reply("ok")
`
)
@@ -3,20 +3,16 @@ package ai360
import (
"github.com/labring/sealos/service/aiproxy/model"
"github.com/labring/sealos/service/aiproxy/relay/adaptor/openai"
"github.com/labring/sealos/service/aiproxy/relay/meta"
)
type Adaptor struct {
openai.Adaptor
}
const baseURL = "https://ai.360.cn"
const baseURL = "https://ai.360.cn/v1"
func (a *Adaptor) GetRequestURL(meta *meta.Meta) (string, error) {
if meta.Channel.BaseURL == "" {
meta.Channel.BaseURL = baseURL
}
return a.Adaptor.GetRequestURL(meta)
func (a *Adaptor) GetBaseURL() string {
return baseURL
}
func (a *Adaptor) GetModelList() []*model.ModelConfig {
+26 -16
View File
@@ -22,6 +22,10 @@ type Adaptor struct{}
const baseURL = "https://dashscope.aliyuncs.com"
func (a *Adaptor) GetBaseURL() string {
return baseURL
}
func (a *Adaptor) GetRequestURL(meta *meta.Meta) (string, error) {
u := meta.Channel.BaseURL
if u == "" {
@@ -34,6 +38,8 @@ func (a *Adaptor) GetRequestURL(meta *meta.Meta) (string, error) {
return u + "/api/v1/services/aigc/text2image/image-synthesis", nil
case relaymode.ChatCompletions:
return u + "/compatible-mode/v1/chat/completions", nil
case relaymode.Completions:
return u + "/compatible-mode/v1/completions", nil
case relaymode.AudioSpeech, relaymode.AudioTranscription:
return u + "/api-ws/v1/inference", nil
case relaymode.Rerank:
@@ -46,13 +52,11 @@ func (a *Adaptor) GetRequestURL(meta *meta.Meta) (string, error) {
func (a *Adaptor) SetupRequestHeader(meta *meta.Meta, _ *gin.Context, req *http.Request) error {
req.Header.Set("Authorization", "Bearer "+meta.Channel.Key)
if meta.Channel.Config.Plugin != "" {
req.Header.Set("X-Dashscope-Plugin", meta.Channel.Config.Plugin)
}
// req.Header.Set("X-Dashscope-Plugin", meta.Channel.Config.Plugin)
return nil
}
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Reader, error) {
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (string, http.Header, io.Reader, error) {
switch meta.Mode {
case relaymode.ImagesGenerations:
return ConvertImageRequest(meta, req)
@@ -60,30 +64,36 @@ func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (http.Heade
return ConvertRerankRequest(meta, req)
case relaymode.Embeddings:
return ConvertEmbeddingsRequest(meta, req)
case relaymode.ChatCompletions:
case relaymode.ChatCompletions, relaymode.Completions:
return openai.ConvertRequest(meta, req)
case relaymode.AudioSpeech:
return ConvertTTSRequest(meta, req)
case relaymode.AudioTranscription:
return ConvertSTTRequest(meta, req)
default:
return nil, nil, errors.New("unsupported convert request mode")
return "", nil, nil, errors.New("unsupported convert request mode")
}
}
func ignoreTest(meta *meta.Meta) bool {
return meta.IsChannelTest &&
(strings.Contains(meta.ActualModel, "-ocr") ||
strings.HasPrefix(meta.ActualModel, "qwen-mt-"))
}
func (a *Adaptor) DoRequest(meta *meta.Meta, _ *gin.Context, req *http.Request) (*http.Response, error) {
if ignoreTest(meta) {
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(bytes.NewReader(nil)),
}, nil
}
switch meta.Mode {
case relaymode.AudioSpeech:
return TTSDoRequest(meta, req)
case relaymode.AudioTranscription:
return STTDoRequest(meta, req)
case relaymode.ChatCompletions:
if meta.IsChannelTest && strings.Contains(meta.ActualModelName, "-ocr") {
return &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(bytes.NewReader(nil)),
}, nil
}
fallthrough
default:
return utils.DoRequest(req)
@@ -91,15 +101,15 @@ func (a *Adaptor) DoRequest(meta *meta.Meta, _ *gin.Context, req *http.Request)
}
func (a *Adaptor) DoResponse(meta *meta.Meta, c *gin.Context, resp *http.Response) (usage *relaymodel.Usage, err *relaymodel.ErrorWithStatusCode) {
if ignoreTest(meta) {
return &relaymodel.Usage{}, nil
}
switch meta.Mode {
case relaymode.Embeddings:
usage, err = EmbeddingsHandler(meta, c, resp)
case relaymode.ImagesGenerations:
usage, err = ImageHandler(meta, c, resp)
case relaymode.ChatCompletions:
if meta.IsChannelTest && strings.Contains(meta.ActualModelName, "-ocr") {
return nil, nil
}
case relaymode.ChatCompletions, relaymode.Completions:
usage, err = openai.DoResponse(meta, c, resp)
case relaymode.Rerank:
usage, err = RerankHandler(meta, c, resp)
File diff suppressed because it is too large Load Diff
@@ -15,16 +15,16 @@ import (
relaymodel "github.com/labring/sealos/service/aiproxy/relay/model"
)
func ConvertEmbeddingsRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Reader, error) {
func ConvertEmbeddingsRequest(meta *meta.Meta, req *http.Request) (string, http.Header, io.Reader, error) {
var reqMap map[string]any
err := common.UnmarshalBodyReusable(req, &reqMap)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
reqMap["model"] = meta.ActualModelName
reqMap["model"] = meta.ActualModel
input, ok := reqMap["input"]
if !ok {
return nil, nil, errors.New("input is required")
return "", nil, nil, errors.New("input is required")
}
switch v := input.(type) {
case string:
@@ -47,16 +47,16 @@ func ConvertEmbeddingsRequest(meta *meta.Meta, req *http.Request) (http.Header,
reqMap["parameters"] = parameters
jsonData, err := json.Marshal(reqMap)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
return nil, bytes.NewReader(jsonData), nil
return http.MethodPost, nil, bytes.NewReader(jsonData), nil
}
func embeddingResponse2OpenAI(meta *meta.Meta, response *EmbeddingResponse) *openai.EmbeddingResponse {
openAIEmbeddingResponse := openai.EmbeddingResponse{
Object: "list",
Data: make([]*openai.EmbeddingResponseItem, 0, 1),
Model: meta.OriginModelName,
Model: meta.OriginModel,
Usage: response.Usage,
}
@@ -94,7 +94,7 @@ func EmbeddingsHandler(meta *meta.Meta, c *gin.Context, resp *http.Response) (*r
}
_, err = c.Writer.Write(data)
if err != nil {
log.Error("write response body failed: " + err.Error())
log.Warnf("write response body failed: %v", err)
}
return &openaiResponse.Usage, nil
}
+7 -8
View File
@@ -11,7 +11,6 @@ import (
"github.com/gin-gonic/gin"
json "github.com/json-iterator/go"
"github.com/labring/sealos/service/aiproxy/common/helper"
"github.com/labring/sealos/service/aiproxy/common/image"
"github.com/labring/sealos/service/aiproxy/middleware"
"github.com/labring/sealos/service/aiproxy/relay/adaptor/openai"
@@ -23,12 +22,12 @@ import (
const MetaResponseFormat = "response_format"
func ConvertImageRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Reader, error) {
func ConvertImageRequest(meta *meta.Meta, req *http.Request) (string, http.Header, io.Reader, error) {
request, err := utils.UnmarshalImageRequest(req)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
request.Model = meta.ActualModelName
request.Model = meta.ActualModel
var imageRequest ImageRequest
imageRequest.Input.Prompt = request.Prompt
@@ -41,9 +40,9 @@ func ConvertImageRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Re
data, err := json.Marshal(&imageRequest)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
return http.Header{
return http.MethodPost, http.Header{
"X-Dashscope-Async": {"enable"},
}, bytes.NewReader(data), nil
}
@@ -98,7 +97,7 @@ func ImageHandler(meta *meta.Meta, c *gin.Context, resp *http.Response) (*model.
c.Writer.WriteHeader(resp.StatusCode)
_, err = c.Writer.Write(jsonResponse)
if err != nil {
log.Error("aliImageHandler write response body failed: " + err.Error())
log.Warnf("aliImageHandler write response body failed: %v", err)
}
return &model.Usage{}, nil
}
@@ -168,7 +167,7 @@ func asyncTaskWait(ctx context.Context, taskID string, key string) (*TaskRespons
func responseAli2OpenAIImage(ctx context.Context, response *TaskResponse, responseFormat string) *openai.ImageResponse {
imageResponse := openai.ImageResponse{
Created: helper.GetTimestamp(),
Created: time.Now().Unix(),
}
for _, data := range response.Output.Results {
+8 -8
View File
@@ -26,13 +26,13 @@ type RerankUsage struct {
TotalTokens int `json:"total_tokens"`
}
func ConvertRerankRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Reader, error) {
func ConvertRerankRequest(meta *meta.Meta, req *http.Request) (string, http.Header, io.Reader, error) {
reqMap := make(map[string]any)
err := common.UnmarshalBodyReusable(req, &reqMap)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
reqMap["model"] = meta.ActualModelName
reqMap["model"] = meta.ActualModel
reqMap["input"] = map[string]any{
"query": reqMap["query"],
"documents": reqMap["documents"],
@@ -50,9 +50,9 @@ func ConvertRerankRequest(meta *meta.Meta, req *http.Request) (http.Header, io.R
reqMap["parameters"] = parameters
jsonData, err := json.Marshal(reqMap)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
return nil, bytes.NewReader(jsonData), nil
return http.MethodPost, nil, bytes.NewReader(jsonData), nil
}
func RerankHandler(meta *meta.Meta, c *gin.Context, resp *http.Response) (*relaymodel.Usage, *relaymodel.ErrorWithStatusCode) {
@@ -86,9 +86,9 @@ func RerankHandler(meta *meta.Meta, c *gin.Context, resp *http.Response) (*relay
var usage *relaymodel.Usage
if rerankResponse.Usage == nil {
usage = &relaymodel.Usage{
PromptTokens: meta.PromptTokens,
PromptTokens: meta.InputTokens,
CompletionTokens: 0,
TotalTokens: meta.PromptTokens,
TotalTokens: meta.InputTokens,
}
} else {
usage = &relaymodel.Usage{
@@ -103,7 +103,7 @@ func RerankHandler(meta *meta.Meta, c *gin.Context, resp *http.Response) (*relay
}
_, err = c.Writer.Write(jsonResponse)
if err != nil {
log.Error("write response body failed: " + err.Error())
log.Warnf("write response body failed: %v", err)
}
return usage, nil
}
@@ -2,7 +2,6 @@ package ali
import (
"bytes"
"errors"
"io"
"net/http"
@@ -59,29 +58,19 @@ type STTUsage struct {
Characters int `json:"characters"`
}
func ConvertSTTRequest(meta *meta.Meta, request *http.Request) (http.Header, io.Reader, error) {
func ConvertSTTRequest(meta *meta.Meta, request *http.Request) (string, http.Header, io.Reader, error) {
err := request.ParseMultipartForm(1024 * 1024 * 4)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
var audioData []byte
if files, ok := request.MultipartForm.File["file"]; !ok {
return nil, nil, errors.New("audio file is required")
} else if len(files) == 1 {
file, err := files[0].Open()
if err != nil {
return nil, nil, err
}
audioData, err = io.ReadAll(file)
file.Close()
if err != nil {
return nil, nil, err
}
} else {
return nil, nil, errors.New("audio file is required")
audioFile, _, err := request.FormFile("file")
if err != nil {
return "", nil, nil, err
}
audioData, err := io.ReadAll(audioFile)
if err != nil {
return "", nil, nil, err
}
sttRequest := STTMessage{
Header: STTHeader{
Action: "run-task",
@@ -89,7 +78,7 @@ func ConvertSTTRequest(meta *meta.Meta, request *http.Request) (http.Header, io.
TaskID: uuid.New().String(),
},
Payload: STTPayload{
Model: meta.ActualModelName,
Model: meta.ActualModel,
Task: "asr",
TaskGroup: "audio",
Function: "recognition",
@@ -103,11 +92,11 @@ func ConvertSTTRequest(meta *meta.Meta, request *http.Request) (http.Header, io.
data, err := json.Marshal(sttRequest)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
meta.Set("audio_data", audioData)
meta.Set("task_id", sttRequest.Header.TaskID)
return http.Header{
return http.MethodPost, http.Header{
"X-DashScope-DataInspection": {"enable"},
}, bytes.NewReader(data), nil
}
+6 -6
View File
@@ -93,21 +93,21 @@ var ttsSupportedFormat = map[string]struct{}{
"mp3": {},
}
func ConvertTTSRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Reader, error) {
func ConvertTTSRequest(meta *meta.Meta, req *http.Request) (string, http.Header, io.Reader, error) {
request, err := utils.UnmarshalTTSRequest(req)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
reqMap, err := utils.UnmarshalMap(req)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
var sampleRate int
sampleRateI, ok := reqMap["sample_rate"].(float64)
if ok {
sampleRate = int(sampleRateI)
}
request.Model = meta.ActualModelName
request.Model = meta.ActualModel
if strings.HasPrefix(request.Model, "sambert-v") {
voice := request.Voice
@@ -156,9 +156,9 @@ func ConvertTTSRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Read
data, err := json.Marshal(ttsRequest)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
return http.Header{
return http.MethodPost, http.Header{
"X-DashScope-DataInspection": {"enable"},
}, bytes.NewReader(data), nil
}
@@ -16,14 +16,14 @@ import (
type Adaptor struct{}
const baseURL = "https://api.anthropic.com"
const baseURL = "https://api.anthropic.com/v1"
func (a *Adaptor) GetBaseURL() string {
return baseURL
}
func (a *Adaptor) GetRequestURL(meta *meta.Meta) (string, error) {
u := meta.Channel.BaseURL
if u == "" {
u = baseURL
}
return u + "/v1/messages", nil
return meta.Channel.BaseURL + "/messages", nil
}
func (a *Adaptor) SetupRequestHeader(meta *meta.Meta, c *gin.Context, req *http.Request) error {
@@ -37,24 +37,24 @@ func (a *Adaptor) SetupRequestHeader(meta *meta.Meta, c *gin.Context, req *http.
// https://x.com/alexalbert__/status/1812921642143900036
// claude-3-5-sonnet can support 8k context
if strings.HasPrefix(meta.ActualModelName, "claude-3-5-sonnet") {
if strings.HasPrefix(meta.ActualModel, "claude-3-5-sonnet") {
req.Header.Set("Anthropic-Beta", "max-tokens-3-5-sonnet-2024-07-15")
}
return nil
}
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Reader, error) {
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (string, http.Header, io.Reader, error) {
data, err := ConvertRequest(meta, req)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
data2, err := json.Marshal(data)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
return nil, bytes.NewReader(data2), nil
return http.MethodPost, nil, bytes.NewReader(data2), nil
}
func (a *Adaptor) DoRequest(_ *meta.Meta, _ *gin.Context, req *http.Request) (*http.Response, error) {
@@ -7,48 +7,73 @@ import (
var ModelList = []*model.ModelConfig{
{
Model: "claude-instant-1.2",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerAnthropic,
Model: "claude-3-haiku-20240307",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerAnthropic,
InputPrice: 0.0025,
OutputPrice: 0.0125,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(200000),
model.WithModelConfigMaxOutputTokens(4096),
),
},
{
Model: "claude-2.0",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerAnthropic,
Model: "claude-3-opus-20240229",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerAnthropic,
InputPrice: 0.015,
OutputPrice: 0.075,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(200000),
model.WithModelConfigMaxOutputTokens(4096),
),
},
{
Model: "claude-2.1",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerAnthropic,
Model: "claude-3-5-haiku-20241022",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerAnthropic,
InputPrice: 0.0008,
OutputPrice: 0.004,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(200000),
model.WithModelConfigMaxOutputTokens(4096),
model.WithModelConfigToolChoice(true),
),
},
{
Model: "claude-3-haiku-20240307",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerAnthropic,
Model: "claude-3-5-sonnet-20240620",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerAnthropic,
InputPrice: 0.003,
OutputPrice: 0.015,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(200000),
model.WithModelConfigMaxOutputTokens(8192),
model.WithModelConfigToolChoice(true),
),
},
{
Model: "claude-3-sonnet-20240229",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerAnthropic,
Model: "claude-3-5-sonnet-20241022",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerAnthropic,
InputPrice: 0.003,
OutputPrice: 0.015,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(200000),
model.WithModelConfigMaxOutputTokens(8192),
model.WithModelConfigToolChoice(true),
),
},
{
Model: "claude-3-opus-20240229",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerAnthropic,
},
{
Model: "claude-3-5-sonnet-20240620",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerAnthropic,
},
{
Model: "claude-3-5-sonnet-20241022",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerAnthropic,
},
{
Model: "claude-3-5-sonnet-latest",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerAnthropic,
Model: "claude-3-5-sonnet-latest",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerAnthropic,
InputPrice: 0.003,
OutputPrice: 0.015,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(200000),
model.WithModelConfigMaxOutputTokens(8192),
model.WithModelConfigToolChoice(true),
),
},
}
+36 -22
View File
@@ -4,17 +4,17 @@ import (
"bufio"
"net/http"
"slices"
json "github.com/json-iterator/go"
"github.com/labring/sealos/service/aiproxy/common/conv"
"github.com/labring/sealos/service/aiproxy/common/render"
"github.com/labring/sealos/service/aiproxy/middleware"
"time"
"github.com/gin-gonic/gin"
json "github.com/json-iterator/go"
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/common/helper"
"github.com/labring/sealos/service/aiproxy/common/conv"
"github.com/labring/sealos/service/aiproxy/common/image"
"github.com/labring/sealos/service/aiproxy/common/render"
"github.com/labring/sealos/service/aiproxy/middleware"
"github.com/labring/sealos/service/aiproxy/relay/adaptor/openai"
"github.com/labring/sealos/service/aiproxy/relay/constant"
"github.com/labring/sealos/service/aiproxy/relay/meta"
"github.com/labring/sealos/service/aiproxy/relay/model"
)
@@ -26,10 +26,8 @@ func stopReasonClaude2OpenAI(reason *string) string {
return ""
}
switch *reason {
case "end_turn":
return "stop"
case "stop_sequence":
return "stop"
case "end_turn", "stop_sequence":
return constant.StopFinishReason
case "max_tokens":
return "length"
case toolUseType:
@@ -45,7 +43,7 @@ func ConvertRequest(meta *meta.Meta, req *http.Request) (*Request, error) {
if err != nil {
return nil, err
}
textRequest.Model = meta.ActualModelName
textRequest.Model = meta.ActualModel
meta.Set("stream", textRequest.Stream)
claudeTools := make([]Tool, 0, len(textRequest.Tools))
@@ -256,18 +254,18 @@ func ResponseClaude2OpenAI(claudeResponse *Response) *openai.TextResponse {
ID: "chatcmpl-" + claudeResponse.ID,
Model: claudeResponse.Model,
Object: "chat.completion",
Created: helper.GetTimestamp(),
Created: time.Now().Unix(),
Choices: []*openai.TextResponseChoice{&choice},
}
return &fullTextResponse
}
func StreamHandler(_ *meta.Meta, c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, *model.Usage) {
func StreamHandler(m *meta.Meta, c *gin.Context, resp *http.Response) (*model.ErrorWithStatusCode, *model.Usage) {
defer resp.Body.Close()
log := middleware.GetLogger(c)
createdTime := helper.GetTimestamp()
createdTime := time.Now().Unix()
scanner := bufio.NewScanner(resp.Body)
scanner.Split(func(data []byte, atEOF bool) (advance int, token []byte, err error) {
if atEOF && len(data) == 0 {
@@ -285,9 +283,9 @@ func StreamHandler(_ *meta.Meta, c *gin.Context, resp *http.Response) (*model.Er
common.SetEventStreamHeaders(c)
var usage model.Usage
var modelName string
var id string
var lastToolCallChoice *openai.ChatCompletionsStreamResponseChoice
var usageWrited bool
for scanner.Scan() {
data := scanner.Bytes()
@@ -314,11 +312,14 @@ func StreamHandler(_ *meta.Meta, c *gin.Context, resp *http.Response) (*model.Er
if meta != nil {
usage.PromptTokens += meta.Usage.InputTokens
usage.CompletionTokens += meta.Usage.OutputTokens
usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens
if len(meta.ID) > 0 { // only message_start has an id, otherwise it's a finish_reason event.
modelName = meta.Model
id = "chatcmpl-" + meta.ID
continue
}
response.Usage = &usage
usageWrited = true
if lastToolCallChoice != nil && len(lastToolCallChoice.Delta.ToolCalls) > 0 {
lastArgs := &lastToolCallChoice.Delta.ToolCalls[len(lastToolCallChoice.Delta.ToolCalls)-1].Function
if len(lastArgs.Arguments) == 0 { // compatible with OpenAI sending an empty object `{}` when no arguments.
@@ -330,7 +331,7 @@ func StreamHandler(_ *meta.Meta, c *gin.Context, resp *http.Response) (*model.Er
}
response.ID = id
response.Model = modelName
response.Model = m.OriginModel
response.Created = createdTime
for _, choice := range response.Choices {
@@ -338,16 +339,29 @@ func StreamHandler(_ *meta.Meta, c *gin.Context, resp *http.Response) (*model.Er
lastToolCallChoice = choice
}
}
err = render.ObjectData(c, response)
if err != nil {
log.Error("error rendering stream response: " + err.Error())
}
_ = render.ObjectData(c, response)
}
if err := scanner.Err(); err != nil {
log.Error("error reading stream: " + err.Error())
}
if usage.CompletionTokens == 0 && usage.PromptTokens == 0 {
usage.PromptTokens = m.InputTokens
}
usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens
if !usageWrited {
_ = render.ObjectData(c, &openai.ChatCompletionsStreamResponse{
Model: m.OriginModel,
Object: "chat.completion.chunk",
Created: createdTime,
Choices: []*openai.ChatCompletionsStreamResponseChoice{},
Usage: &usage,
})
}
render.Done(c)
return nil, &usage
@@ -373,7 +387,7 @@ func Handler(meta *meta.Meta, c *gin.Context, resp *http.Response) (*model.Error
}, nil
}
fullTextResponse := ResponseClaude2OpenAI(&claudeResponse)
fullTextResponse.Model = meta.OriginModelName
fullTextResponse.Model = meta.OriginModel
usage := model.Usage{
PromptTokens: claudeResponse.Usage.InputTokens,
CompletionTokens: claudeResponse.Usage.OutputTokens,
+7 -3
View File
@@ -17,10 +17,14 @@ var _ adaptor.Adaptor = new(Adaptor)
type Adaptor struct{}
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Reader, error) {
adaptor := GetAdaptor(meta.ActualModelName)
func (a *Adaptor) GetBaseURL() string {
return ""
}
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (string, http.Header, io.Reader, error) {
adaptor := GetAdaptor(meta.ActualModel)
if adaptor == nil {
return nil, nil, errors.New("adaptor not found")
return "", nil, nil, errors.New("adaptor not found")
}
meta.Set("awsAdapter", adaptor)
return adaptor.ConvertRequest(meta, req)
@@ -19,13 +19,13 @@ var _ utils.AwsAdapter = new(Adaptor)
type Adaptor struct{}
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Reader, error) {
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (string, http.Header, io.Reader, error) {
r, err := anthropic.ConvertRequest(meta, req)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
meta.Set(ConvertedRequest, r)
return nil, nil, nil
return "", nil, nil, nil
}
func (a *Adaptor) DoResponse(meta *meta.Meta, c *gin.Context) (usage *model.Usage, err *model.ErrorWithStatusCode) {
@@ -4,6 +4,7 @@ package aws
import (
"io"
"net/http"
"time"
json "github.com/json-iterator/go"
@@ -12,7 +13,6 @@ import (
"github.com/aws/aws-sdk-go-v2/service/bedrockruntime/types"
"github.com/gin-gonic/gin"
"github.com/jinzhu/copier"
"github.com/labring/sealos/service/aiproxy/common/helper"
"github.com/labring/sealos/service/aiproxy/common/render"
"github.com/labring/sealos/service/aiproxy/middleware"
"github.com/labring/sealos/service/aiproxy/model"
@@ -93,7 +93,7 @@ func awsModelID(requestModel string) (string, error) {
}
func Handler(meta *meta.Meta, c *gin.Context) (*relaymodel.ErrorWithStatusCode, *relaymodel.Usage) {
awsModelID, err := awsModelID(meta.ActualModelName)
awsModelID, err := awsModelID(meta.ActualModel)
if err != nil {
return utils.WrapErr(errors.Wrap(err, "awsModelID")), nil
}
@@ -121,7 +121,12 @@ func Handler(meta *meta.Meta, c *gin.Context) (*relaymodel.ErrorWithStatusCode,
return utils.WrapErr(errors.Wrap(err, "marshal request")), nil
}
awsResp, err := meta.AwsClient().InvokeModel(c.Request.Context(), awsReq)
awsClient, err := utils.AwsClientFromMeta(meta)
if err != nil {
return utils.WrapErr(errors.Wrap(err, "get aws client")), nil
}
awsResp, err := awsClient.InvokeModel(c.Request.Context(), awsReq)
if err != nil {
return utils.WrapErr(errors.Wrap(err, "InvokeModel")), nil
}
@@ -133,7 +138,7 @@ func Handler(meta *meta.Meta, c *gin.Context) (*relaymodel.ErrorWithStatusCode,
}
openaiResp := anthropic.ResponseClaude2OpenAI(claudeResponse)
openaiResp.Model = meta.OriginModelName
openaiResp.Model = meta.OriginModel
usage := relaymodel.Usage{
PromptTokens: claudeResponse.Usage.InputTokens,
CompletionTokens: claudeResponse.Usage.OutputTokens,
@@ -147,9 +152,9 @@ func Handler(meta *meta.Meta, c *gin.Context) (*relaymodel.ErrorWithStatusCode,
func StreamHandler(meta *meta.Meta, c *gin.Context) (*relaymodel.ErrorWithStatusCode, *relaymodel.Usage) {
log := middleware.GetLogger(c)
createdTime := helper.GetTimestamp()
originModelName := meta.OriginModelName
awsModelID, err := awsModelID(meta.ActualModelName)
createdTime := time.Now().Unix()
originModelName := meta.OriginModel
awsModelID, err := awsModelID(meta.ActualModel)
if err != nil {
return utils.WrapErr(errors.Wrap(err, "awsModelID")), nil
}
@@ -177,7 +182,12 @@ func StreamHandler(meta *meta.Meta, c *gin.Context) (*relaymodel.ErrorWithStatus
return utils.WrapErr(errors.Wrap(err, "marshal request")), nil
}
awsResp, err := meta.AwsClient().InvokeModelWithResponseStream(c.Request.Context(), awsReq)
awsClient, err := utils.AwsClientFromMeta(meta)
if err != nil {
return utils.WrapErr(errors.Wrap(err, "get aws client")), nil
}
awsResp, err := awsClient.InvokeModelWithResponseStream(c.Request.Context(), awsReq)
if err != nil {
return utils.WrapErr(errors.Wrap(err, "InvokeModelWithResponseStream")), nil
}
+20
View File
@@ -0,0 +1,20 @@
package aws
import (
"github.com/labring/sealos/service/aiproxy/relay/adaptor"
"github.com/labring/sealos/service/aiproxy/relay/adaptor/aws/utils"
)
var _ adaptor.KeyValidator = (*Adaptor)(nil)
func (a *Adaptor) ValidateKey(key string) error {
_, err := utils.GetAwsConfigFromKey(key)
if err != nil {
return err
}
return nil
}
func (a *Adaptor) KeyHelp() string {
return "region|ak|sk"
}
@@ -19,16 +19,16 @@ var _ utils.AwsAdapter = new(Adaptor)
type Adaptor struct{}
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Reader, error) {
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (string, http.Header, io.Reader, error) {
request, err := relayutils.UnmarshalGeneralOpenAIRequest(req)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
request.Model = meta.ActualModelName
request.Model = meta.ActualModel
meta.Set("stream", request.Stream)
llamaReq := ConvertRequest(request)
meta.Set(ConvertedRequest, llamaReq)
return nil, nil, nil
return "", nil, nil, nil
}
func (a *Adaptor) DoResponse(meta *meta.Meta, c *gin.Context) (usage *model.Usage, err *model.ErrorWithStatusCode) {
@@ -6,22 +6,22 @@ import (
"io"
"net/http"
"text/template"
"time"
json "github.com/json-iterator/go"
"github.com/labring/sealos/service/aiproxy/common/random"
"github.com/labring/sealos/service/aiproxy/common/render"
"github.com/labring/sealos/service/aiproxy/middleware"
"github.com/labring/sealos/service/aiproxy/model"
"github.com/aws/aws-sdk-go-v2/aws"
"github.com/aws/aws-sdk-go-v2/service/bedrockruntime"
"github.com/aws/aws-sdk-go-v2/service/bedrockruntime/types"
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/common/helper"
"github.com/labring/sealos/service/aiproxy/common/random"
"github.com/labring/sealos/service/aiproxy/common/render"
"github.com/labring/sealos/service/aiproxy/middleware"
"github.com/labring/sealos/service/aiproxy/model"
"github.com/labring/sealos/service/aiproxy/relay/adaptor/aws/utils"
"github.com/labring/sealos/service/aiproxy/relay/adaptor/openai"
"github.com/labring/sealos/service/aiproxy/relay/constant"
"github.com/labring/sealos/service/aiproxy/relay/meta"
relaymodel "github.com/labring/sealos/service/aiproxy/relay/model"
"github.com/labring/sealos/service/aiproxy/relay/relaymode"
@@ -94,7 +94,7 @@ func ConvertRequest(textRequest *relaymodel.GeneralOpenAIRequest) *Request {
}
func Handler(meta *meta.Meta, c *gin.Context) (*relaymodel.ErrorWithStatusCode, *relaymodel.Usage) {
awsModelID, err := awsModelID(meta.ActualModelName)
awsModelID, err := awsModelID(meta.ActualModel)
if err != nil {
return utils.WrapErr(errors.Wrap(err, "awsModelID")), nil
}
@@ -115,7 +115,12 @@ func Handler(meta *meta.Meta, c *gin.Context) (*relaymodel.ErrorWithStatusCode,
return utils.WrapErr(errors.Wrap(err, "marshal request")), nil
}
awsResp, err := meta.AwsClient().InvokeModel(c.Request.Context(), awsReq)
awsClient, err := utils.AwsClientFromMeta(meta)
if err != nil {
return utils.WrapErr(errors.Wrap(err, "get aws client")), nil
}
awsResp, err := awsClient.InvokeModel(c.Request.Context(), awsReq)
if err != nil {
return utils.WrapErr(errors.Wrap(err, "InvokeModel")), nil
}
@@ -127,7 +132,7 @@ func Handler(meta *meta.Meta, c *gin.Context) (*relaymodel.ErrorWithStatusCode,
}
openaiResp := ResponseLlama2OpenAI(&llamaResponse)
openaiResp.Model = meta.OriginModelName
openaiResp.Model = meta.OriginModel
usage := relaymodel.Usage{
PromptTokens: llamaResponse.PromptTokenCount,
CompletionTokens: llamaResponse.GenerationTokenCount,
@@ -156,7 +161,7 @@ func ResponseLlama2OpenAI(llamaResponse *Response) *openai.TextResponse {
fullTextResponse := openai.TextResponse{
ID: "chatcmpl-" + random.GetUUID(),
Object: "chat.completion",
Created: helper.GetTimestamp(),
Created: time.Now().Unix(),
Choices: []*openai.TextResponseChoice{&choice},
}
return &fullTextResponse
@@ -165,8 +170,8 @@ func ResponseLlama2OpenAI(llamaResponse *Response) *openai.TextResponse {
func StreamHandler(meta *meta.Meta, c *gin.Context) (*relaymodel.ErrorWithStatusCode, *relaymodel.Usage) {
log := middleware.GetLogger(c)
createdTime := helper.GetTimestamp()
awsModelID, err := awsModelID(meta.ActualModelName)
createdTime := time.Now().Unix()
awsModelID, err := awsModelID(meta.ActualModel)
if err != nil {
return utils.WrapErr(errors.Wrap(err, "awsModelID")), nil
}
@@ -187,7 +192,12 @@ func StreamHandler(meta *meta.Meta, c *gin.Context) (*relaymodel.ErrorWithStatus
return utils.WrapErr(errors.Wrap(err, "marshal request")), nil
}
awsResp, err := meta.AwsClient().InvokeModelWithResponseStream(c.Request.Context(), awsReq)
awsClient, err := utils.AwsClientFromMeta(meta)
if err != nil {
return utils.WrapErr(errors.Wrap(err, "get aws client")), nil
}
awsResp, err := awsClient.InvokeModelWithResponseStream(c.Request.Context(), awsReq)
if err != nil {
return utils.WrapErr(errors.Wrap(err, "InvokeModelWithResponseStream")), nil
}
@@ -215,13 +225,13 @@ func StreamHandler(meta *meta.Meta, c *gin.Context) (*relaymodel.ErrorWithStatus
if llamaResp.PromptTokenCount > 0 {
usage.PromptTokens = llamaResp.PromptTokenCount
}
if llamaResp.StopReason == "stop" {
if llamaResp.StopReason == constant.StopFinishReason {
usage.CompletionTokens = llamaResp.GenerationTokenCount
usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens
}
response := StreamResponseLlama2OpenAI(&llamaResp)
response.ID = "chatcmpl-" + random.GetUUID()
response.Model = meta.OriginModelName
response.Model = meta.OriginModel
response.Created = createdTime
err = render.ObjectData(c, response)
if err != nil {
@@ -1,15 +1,68 @@
package utils
import (
"errors"
"io"
"net/http"
"strings"
"github.com/aws/aws-sdk-go-v2/aws"
"github.com/aws/aws-sdk-go-v2/credentials"
"github.com/aws/aws-sdk-go-v2/service/bedrockruntime"
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/relay/meta"
"github.com/labring/sealos/service/aiproxy/relay/model"
)
type AwsAdapter interface {
ConvertRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Reader, error)
ConvertRequest(meta *meta.Meta, req *http.Request) (string, http.Header, io.Reader, error)
DoResponse(meta *meta.Meta, c *gin.Context) (usage *model.Usage, err *model.ErrorWithStatusCode)
}
type AwsConfig struct {
Region string
AK string
SK string
}
func GetAwsConfigFromKey(key string) (*AwsConfig, error) {
split := strings.Split(key, "|")
if len(split) != 3 {
return nil, errors.New("invalid key format")
}
return &AwsConfig{
Region: split[0],
AK: split[1],
SK: split[2],
}, nil
}
func AwsClient(config *AwsConfig) *bedrockruntime.Client {
return bedrockruntime.New(bedrockruntime.Options{
Region: config.Region,
Credentials: aws.NewCredentialsCache(credentials.NewStaticCredentialsProvider(config.AK, config.SK, "")),
})
}
func awsClientFromKey(key string) (*bedrockruntime.Client, error) {
config, err := GetAwsConfigFromKey(key)
if err != nil {
return nil, err
}
return AwsClient(config), nil
}
const AwsClientKey = "aws_client"
func AwsClientFromMeta(meta *meta.Meta) (*bedrockruntime.Client, error) {
awsClientI, ok := meta.Get(AwsClientKey)
if ok {
return awsClientI.(*bedrockruntime.Client), nil
}
awsClient, err := awsClientFromKey(meta.Channel.Key)
if err != nil {
return nil, err
}
meta.Set(AwsClientKey, awsClient)
return awsClient, nil
}
@@ -8,48 +8,48 @@ import (
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/relay/adaptor/openai"
"github.com/labring/sealos/service/aiproxy/relay/meta"
"github.com/labring/sealos/service/aiproxy/relay/relaymode"
)
type Adaptor struct {
openai.Adaptor
}
// func (a *Adaptor) GetRequestURL(meta *meta.Meta) (string, error) {
// switch meta.Mode {
// case relaymode.ImagesGenerations:
// // https://learn.microsoft.com/en-us/azure/ai-services/openai/dall-e-quickstart?tabs=dalle3%2Ccommand-line&pivots=rest-api
// // https://{resource_name}.openai.azure.com/openai/deployments/dall-e-3/images/generations?api-version=2024-03-01-preview
// return fmt.Sprintf("%s/openai/deployments/%s/images/generations?api-version=%s", meta.Channel.BaseURL, meta.ActualModelName, meta.Channel.Config.APIVersion), nil
// case relaymode.AudioTranscription:
// // https://learn.microsoft.com/en-us/azure/ai-services/openai/whisper-quickstart?tabs=command-line#rest-api
// return fmt.Sprintf("%s/openai/deployments/%s/audio/transcriptions?api-version=%s", meta.Channel.BaseURL, meta.ActualModelName, meta.Channel.Config.APIVersion), nil
// case relaymode.AudioSpeech:
// // https://learn.microsoft.com/en-us/azure/ai-services/openai/text-to-speech-quickstart?tabs=command-line#rest-api
// return fmt.Sprintf("%s/openai/deployments/%s/audio/speech?api-version=%s", meta.Channel.BaseURL, meta.ActualModelName, meta.Channel.Config.APIVersion), nil
// }
func (a *Adaptor) GetBaseURL() string {
return "https://{resource_name}.openai.azure.com"
}
// // https://learn.microsoft.com/en-us/azure/cognitive-services/openai/chatgpt-quickstart?pivots=rest-api&tabs=command-line#rest-api
// requestURL := strings.Split(meta.RequestURLPath, "?")[0]
// requestURL = fmt.Sprintf("%s?api-version=%s", requestURL, meta.Channel.Config.APIVersion)
// task := strings.TrimPrefix(requestURL, "/v1/")
// model := strings.ReplaceAll(meta.ActualModelName, ".", "")
// // https://github.com/labring/sealos/service/aiproxy/issues/1191
// // {your endpoint}/openai/deployments/{your azure_model}/chat/completions?api-version={api_version}
// requestURL = fmt.Sprintf("/openai/deployments/%s/%s", model, task)
// return GetFullRequestURL(meta.Channel.BaseURL, requestURL), nil
// }
func GetFullRequestURL(baseURL string, requestURL string) string {
fullRequestURL := fmt.Sprintf("%s%s", baseURL, requestURL)
if strings.HasPrefix(baseURL, "https://gateway.ai.cloudflare.com") {
fullRequestURL = fmt.Sprintf("%s%s", baseURL, strings.TrimPrefix(requestURL, "/v1"))
func (a *Adaptor) GetRequestURL(meta *meta.Meta) (string, error) {
_, apiVersion, err := getTokenAndAPIVersion(meta.Channel.Key)
if err != nil {
return "", err
}
model := strings.ReplaceAll(meta.ActualModel, ".", "")
switch meta.Mode {
case relaymode.ImagesGenerations:
// https://learn.microsoft.com/en-us/azure/ai-services/openai/dall-e-quickstart?tabs=dalle3%2Ccommand-line&pivots=rest-api
// https://{resource_name}.openai.azure.com/openai/deployments/dall-e-3/images/generations?api-version=2024-03-01-preview
return fmt.Sprintf("%s/openai/deployments/%s/images/generations?api-version=%s", meta.Channel.BaseURL, model, apiVersion), nil
case relaymode.AudioTranscription:
// https://learn.microsoft.com/en-us/azure/ai-services/openai/whisper-quickstart?tabs=command-line#rest-api
return fmt.Sprintf("%s/openai/deployments/%s/audio/transcriptions?api-version=%s", meta.Channel.BaseURL, model, apiVersion), nil
case relaymode.AudioSpeech:
// https://learn.microsoft.com/en-us/azure/ai-services/openai/text-to-speech-quickstart?tabs=command-line#rest-api
return fmt.Sprintf("%s/openai/deployments/%s/audio/speech?api-version=%s", meta.Channel.BaseURL, model, apiVersion), nil
case relaymode.ChatCompletions:
// https://learn.microsoft.com/en-us/azure/cognitive-services/openai/chatgpt-quickstart?pivots=rest-api&tabs=command-line#rest-api
return fmt.Sprintf("%s/openai/deployments/%s/chat/completions?api-version=%s", meta.Channel.BaseURL, model, apiVersion), nil
default:
return "", fmt.Errorf("unsupported mode: %d", meta.Mode)
}
return fullRequestURL
}
func (a *Adaptor) SetupRequestHeader(meta *meta.Meta, _ *gin.Context, req *http.Request) error {
req.Header.Set("Api-Key", meta.Channel.Key)
token, _, err := getTokenAndAPIVersion(meta.Channel.Key)
if err != nil {
return err
}
req.Header.Set("Api-Key", token)
return nil
}
@@ -0,0 +1,33 @@
package azure
import (
"errors"
"strings"
"github.com/labring/sealos/service/aiproxy/relay/adaptor"
)
var _ adaptor.KeyValidator = (*Adaptor)(nil)
func (a *Adaptor) ValidateKey(key string) error {
_, _, err := getTokenAndAPIVersion(key)
if err != nil {
return err
}
return nil
}
func (a *Adaptor) KeyHelp() string {
return "key or key|api-version"
}
func getTokenAndAPIVersion(key string) (string, string, error) {
split := strings.Split(key, "|")
if len(split) == 1 {
return key, "", nil
}
if len(split) != 2 {
return "", "", errors.New("invalid key format")
}
return split[0], split[1], nil
}
@@ -3,20 +3,16 @@ package baichuan
import (
"github.com/labring/sealos/service/aiproxy/model"
"github.com/labring/sealos/service/aiproxy/relay/adaptor/openai"
"github.com/labring/sealos/service/aiproxy/relay/meta"
)
type Adaptor struct {
openai.Adaptor
}
const baseURL = "https://api.baichuan-ai.com"
const baseURL = "https://api.baichuan-ai.com/v1"
func (a *Adaptor) GetRequestURL(meta *meta.Meta) (string, error) {
if meta.Channel.BaseURL == "" {
meta.Channel.BaseURL = baseURL
}
return a.Adaptor.GetRequestURL(meta)
func (a *Adaptor) GetBaseURL() string {
return baseURL
}
func (a *Adaptor) GetModelList() []*model.ModelConfig {
@@ -7,18 +7,63 @@ import (
var ModelList = []*model.ModelConfig{
{
Model: "Baichuan2-Turbo",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaichuan,
Model: "Baichuan4-Turbo",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaichuan,
InputPrice: 0.015,
OutputPrice: 0.015,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(32768),
),
},
{
Model: "Baichuan2-Turbo-192k",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaichuan,
Model: "Baichuan4-Air",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaichuan,
InputPrice: 0.00098,
OutputPrice: 0.00098,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(32768),
),
},
{
Model: "Baichuan-Text-Embedding",
Type: relaymode.Embeddings,
Owner: model.ModelOwnerBaichuan,
Model: "Baichuan4",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaichuan,
InputPrice: 0.1,
OutputPrice: 0.1,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(32768),
),
},
{
Model: "Baichuan3-Turbo",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaichuan,
InputPrice: 0.012,
OutputPrice: 0.012,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(32768),
),
},
{
Model: "Baichuan3-Turbo-128k",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaichuan,
InputPrice: 0.024,
OutputPrice: 0.024,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(131072),
),
},
{
Model: "Baichuan-Text-Embedding",
Type: relaymode.Embeddings,
Owner: model.ModelOwnerBaichuan,
InputPrice: 0.0005,
Config: model.NewModelConfig(
model.WithModelConfigMaxInputTokens(512),
),
},
}
@@ -23,6 +23,10 @@ const (
baseURL = "https://aip.baidubce.com"
)
func (a *Adaptor) GetBaseURL() string {
return baseURL
}
// Get model-specific endpoint using map
var modelEndpointMap = map[string]string{
"ERNIE-4.0-8K": "completions_pro",
@@ -64,9 +68,9 @@ func (a *Adaptor) GetRequestURL(meta *meta.Meta) (string, error) {
pathSuffix = "text2image"
}
modelEndpoint, ok := modelEndpointMap[meta.ActualModelName]
modelEndpoint, ok := modelEndpointMap[meta.ActualModel]
if !ok {
modelEndpoint = strings.ToLower(meta.ActualModelName)
modelEndpoint = strings.ToLower(meta.ActualModel)
}
// Construct full URL
@@ -86,7 +90,7 @@ func (a *Adaptor) SetupRequestHeader(meta *meta.Meta, _ *gin.Context, req *http.
return nil
}
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Reader, error) {
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (string, http.Header, io.Reader, error) {
switch meta.Mode {
case relaymode.Embeddings:
meta.Set(openai.MetaEmbeddingsPatchInputToSlices, true)
@@ -12,9 +12,9 @@ var ModelList = []*model.ModelConfig{
Owner: model.ModelOwnerBaidu,
InputPrice: 0.004,
OutputPrice: 0.004,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 4800,
},
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(4800),
),
},
{
@@ -23,6 +23,7 @@ var ModelList = []*model.ModelConfig{
Owner: model.ModelOwnerBaidu,
InputPrice: 0.0005,
OutputPrice: 0,
RPM: 1200,
},
{
Model: "bge-large-zh",
@@ -30,6 +31,7 @@ var ModelList = []*model.ModelConfig{
Owner: model.ModelOwnerBAAI,
InputPrice: 0.0005,
OutputPrice: 0,
RPM: 1200,
},
{
Model: "bge-large-en",
@@ -37,6 +39,7 @@ var ModelList = []*model.ModelConfig{
Owner: model.ModelOwnerBAAI,
InputPrice: 0.0005,
OutputPrice: 0,
RPM: 1200,
},
{
Model: "tao-8k",
@@ -44,6 +47,7 @@ var ModelList = []*model.ModelConfig{
Owner: model.ModelOwnerBaidu,
InputPrice: 0.0005,
OutputPrice: 0,
RPM: 1200,
},
{
@@ -52,6 +56,7 @@ var ModelList = []*model.ModelConfig{
Owner: model.ModelOwnerBaidu,
InputPrice: 0.0005,
OutputPrice: 0,
RPM: 1200,
},
{
@@ -41,7 +41,7 @@ func EmbeddingsHandler(meta *meta.Meta, c *gin.Context, resp *http.Response) (*r
if err != nil {
return &baiduResponse.Usage, openai.ErrorWrapper(err, "unmarshal_response_body_failed", http.StatusInternalServerError)
}
respMap["model"] = meta.OriginModelName
respMap["model"] = meta.OriginModel
respMap["object"] = "list"
data, err := json.Marshal(respMap)
@@ -50,7 +50,7 @@ func EmbeddingsHandler(meta *meta.Meta, c *gin.Context, resp *http.Response) (*r
}
_, err = c.Writer.Write(data)
if err != nil {
log.Error("write response body failed: " + err.Error())
log.Warnf("write response body failed: %v", err)
}
return &baiduResponse.Usage, nil
}
+1 -1
View File
@@ -55,7 +55,7 @@ func ImageHandler(_ *meta.Meta, c *gin.Context, resp *http.Response) (*model.Usa
}
_, err = c.Writer.Write(data)
if err != nil {
log.Error("write response body failed: " + err.Error())
log.Warnf("write response body failed: %v", err)
}
return usage, nil
}
+9 -12
View File
@@ -41,12 +41,12 @@ type ChatRequest struct {
EnableCitation bool `json:"enable_citation,omitempty"`
}
func ConvertRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Reader, error) {
func ConvertRequest(meta *meta.Meta, req *http.Request) (string, http.Header, io.Reader, error) {
request, err := utils.UnmarshalGeneralOpenAIRequest(req)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
request.Model = meta.ActualModelName
request.Model = meta.ActualModel
baiduRequest := ChatRequest{
Messages: request.Messages,
Temperature: request.Temperature,
@@ -81,9 +81,9 @@ func ConvertRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Reader,
data, err := json.Marshal(baiduRequest)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
return nil, bytes.NewReader(data), nil
return http.MethodPost, nil, bytes.NewReader(data), nil
}
func responseBaidu2OpenAI(response *ChatResponse) *openai.TextResponse {
@@ -93,7 +93,7 @@ func responseBaidu2OpenAI(response *ChatResponse) *openai.TextResponse {
Role: "assistant",
Content: response.Result,
},
FinishReason: "stop",
FinishReason: constant.StopFinishReason,
}
fullTextResponse := openai.TextResponse{
ID: response.ID,
@@ -117,7 +117,7 @@ func streamResponseBaidu2OpenAI(meta *meta.Meta, baiduResponse *ChatStreamRespon
ID: baiduResponse.ID,
Object: "chat.completion.chunk",
Created: baiduResponse.Created,
Model: meta.OriginModelName,
Model: meta.OriginModel,
Choices: []*openai.ChatCompletionsStreamResponseChoice{&choice},
Usage: baiduResponse.Usage,
}
@@ -158,10 +158,7 @@ func StreamHandler(meta *meta.Meta, c *gin.Context, resp *http.Response) (*model
usage.CompletionTokens = baiduResponse.Usage.TotalTokens - baiduResponse.Usage.PromptTokens
}
response := streamResponseBaidu2OpenAI(meta, &baiduResponse)
err = render.ObjectData(c, response)
if err != nil {
log.Error("error rendering stream response: " + err.Error())
}
_ = render.ObjectData(c, response)
}
if err := scanner.Err(); err != nil {
@@ -185,7 +182,7 @@ func Handler(meta *meta.Meta, c *gin.Context, resp *http.Response) (*model.Usage
return nil, openai.ErrorWrapperWithMessage(baiduResponse.Error.ErrorMsg, "baidu_error_"+strconv.Itoa(baiduResponse.Error.ErrorCode), http.StatusInternalServerError)
}
fullTextResponse := responseBaidu2OpenAI(&baiduResponse)
fullTextResponse.Model = meta.OriginModelName
fullTextResponse.Model = meta.OriginModel
jsonResponse, err := json.Marshal(fullTextResponse)
if err != nil {
return nil, openai.ErrorWrapper(err, "marshal_response_body_failed", http.StatusInternalServerError)
@@ -56,7 +56,7 @@ func RerankHandler(_ *meta.Meta, c *gin.Context, resp *http.Response) (*model.Us
}
_, err = c.Writer.Write(jsonData)
if err != nil {
log.Error("write response body failed: " + err.Error())
log.Warnf("write response body failed: %v", err)
}
return &reRankResp.Usage, nil
}
+2 -2
View File
@@ -10,7 +10,7 @@ import (
"time"
json "github.com/json-iterator/go"
"github.com/labring/sealos/service/aiproxy/common/client"
"github.com/labring/sealos/service/aiproxy/relay/utils"
log "github.com/sirupsen/logrus"
)
@@ -66,7 +66,7 @@ func getBaiduAccessTokenHelper(ctx context.Context, apiKey string) (*AccessToken
}
req.Header.Add("Content-Type", "application/json")
req.Header.Add("Accept", "application/json")
res, err := client.ImpatientHTTPClient.Do(req)
res, err := utils.DoRequest(req)
if err != nil {
return nil, err
}
@@ -7,43 +7,29 @@ import (
"net/http"
"strings"
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/model"
"github.com/labring/sealos/service/aiproxy/relay/adaptor/openai"
"github.com/labring/sealos/service/aiproxy/relay/meta"
relaymodel "github.com/labring/sealos/service/aiproxy/relay/model"
"github.com/labring/sealos/service/aiproxy/relay/relaymode"
"github.com/labring/sealos/service/aiproxy/relay/utils"
"github.com/gin-gonic/gin"
relaymodel "github.com/labring/sealos/service/aiproxy/relay/model"
)
type Adaptor struct{}
const (
baseURL = "https://qianfan.baidubce.com"
baseURL = "https://qianfan.baidubce.com/v2"
)
func (a *Adaptor) GetBaseURL() string {
return baseURL
}
// https://cloud.baidu.com/doc/WENXINWORKSHOP/s/Fm2vrveyu
var v2ModelMap = map[string]string{
"ERNIE-4.0-8K-Latest": "ernie-4.0-8k-latest",
"ERNIE-4.0-8K-Preview": "ernie-4.0-8k-preview",
"ERNIE-4.0-8K": "ernie-4.0-8k",
"ERNIE-4.0-Turbo-8K-Latest": "ernie-4.0-turbo-8k-latest",
"ERNIE-4.0-Turbo-8K-Preview": "ernie-4.0-turbo-8k-preview",
"ERNIE-4.0-Turbo-8K": "ernie-4.0-turbo-8k",
"ERNIE-4.0-Turbo-128K": "ernie-4.0-turbo-128k",
"ERNIE-3.5-8K-Preview": "ernie-3.5-8k-preview",
"ERNIE-3.5-8K": "ernie-3.5-8k",
"ERNIE-3.5-128K": "ernie-3.5-128k",
"ERNIE-Speed-8K": "ernie-speed-8k",
"ERNIE-Speed-128K": "ernie-speed-128k",
"ERNIE-Speed-Pro-128K": "ernie-speed-pro-128k",
"ERNIE-Lite-8K": "ernie-lite-8k",
"ERNIE-Lite-Pro-128K": "ernie-lite-pro-128k",
"ERNIE-Tiny-8K": "ernie-tiny-8k",
"ERNIE-Character-8K": "ernie-char-8k",
"ERNIE-Character-Fiction-8K": "ernie-char-fiction-8k",
"ERNIE-Novel-8K": "ernie-novel-8k",
}
func toV2ModelName(modelName string) string {
@@ -54,13 +40,9 @@ func toV2ModelName(modelName string) string {
}
func (a *Adaptor) GetRequestURL(meta *meta.Meta) (string, error) {
if meta.Channel.BaseURL == "" {
meta.Channel.BaseURL = baseURL
}
switch meta.Mode {
case relaymode.ChatCompletions:
return meta.Channel.BaseURL + "/v2/chat/completions", nil
return meta.Channel.BaseURL + "/chat/completions", nil
default:
return "", fmt.Errorf("unsupported mode: %d", meta.Mode)
}
@@ -75,16 +57,18 @@ func (a *Adaptor) SetupRequestHeader(meta *meta.Meta, _ *gin.Context, req *http.
return nil
}
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Reader, error) {
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (string, http.Header, io.Reader, error) {
switch meta.Mode {
case relaymode.ChatCompletions:
actModel := meta.ActualModelName
actModel := meta.ActualModel
v2Model := toV2ModelName(actModel)
meta.ActualModelName = v2Model
defer func() { meta.ActualModelName = actModel }()
if v2Model != actModel {
meta.ActualModel = v2Model
defer func() { meta.ActualModel = actModel }()
}
return openai.ConvertRequest(meta, req)
default:
return nil, nil, fmt.Errorf("unsupported mode: %d", meta.Mode)
return "", nil, nil, fmt.Errorf("unsupported mode: %d", meta.Mode)
}
}
+166 -115
View File
@@ -5,18 +5,36 @@ import (
"github.com/labring/sealos/service/aiproxy/relay/relaymode"
)
// https://cloud.baidu.com/doc/WENXINWORKSHOP/s/Fm2vrveyu
var ModelList = []*model.ModelConfig{
{
Model: "ERNIE-4.0-8K-Latest",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaidu,
InputPrice: 0.03,
OutputPrice: 0.09,
RPM: 120,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(5120),
model.WithModelConfigMaxInputTokens(5120),
model.WithModelConfigMaxOutputTokens(2048),
model.WithModelConfigToolChoice(true),
),
},
{
Model: "ERNIE-4.0-8K-Preview",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaidu,
InputPrice: 0.03,
OutputPrice: 0.09,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 5120,
model.ModelConfigMaxInputTokensKey: 5120,
model.ModelConfigMaxOutputTokensKey: 2048,
},
RPM: 300,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(5120),
model.WithModelConfigMaxInputTokens(5120),
model.WithModelConfigMaxOutputTokens(2048),
model.WithModelConfigToolChoice(true),
),
},
{
Model: "ERNIE-4.0-8K",
@@ -24,35 +42,13 @@ var ModelList = []*model.ModelConfig{
Owner: model.ModelOwnerBaidu,
InputPrice: 0.03,
OutputPrice: 0.09,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 5120,
model.ModelConfigMaxInputTokensKey: 5120,
model.ModelConfigMaxOutputTokensKey: 2048,
},
},
{
Model: "ERNIE-4.0-8K-Latest",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaidu,
InputPrice: 0.03,
OutputPrice: 0.09,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 5120,
model.ModelConfigMaxInputTokensKey: 5120,
model.ModelConfigMaxOutputTokensKey: 2048,
},
},
{
Model: "ERNIE-4.0-Turbo-8K",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaidu,
InputPrice: 0.02,
OutputPrice: 0.06,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 6144,
model.ModelConfigMaxInputTokensKey: 6144,
model.ModelConfigMaxOutputTokensKey: 2048,
},
RPM: 10000,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(5120),
model.WithModelConfigMaxInputTokens(5120),
model.WithModelConfigMaxOutputTokens(2048),
model.WithModelConfigToolChoice(true),
),
},
{
Model: "ERNIE-4.0-Turbo-8K-Latest",
@@ -60,11 +56,13 @@ var ModelList = []*model.ModelConfig{
Owner: model.ModelOwnerBaidu,
InputPrice: 0.02,
OutputPrice: 0.06,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 6144,
model.ModelConfigMaxInputTokensKey: 6144,
model.ModelConfigMaxOutputTokensKey: 2048,
},
RPM: 60,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(6144),
model.WithModelConfigMaxInputTokens(6144),
model.WithModelConfigMaxOutputTokens(2048),
model.WithModelConfigToolChoice(true),
),
},
{
Model: "ERNIE-4.0-Turbo-8K-Preview",
@@ -72,11 +70,27 @@ var ModelList = []*model.ModelConfig{
Owner: model.ModelOwnerBaidu,
InputPrice: 0.02,
OutputPrice: 0.06,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 6144,
model.ModelConfigMaxInputTokensKey: 6144,
model.ModelConfigMaxOutputTokensKey: 2048,
},
RPM: 60,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(6144),
model.WithModelConfigMaxInputTokens(6144),
model.WithModelConfigMaxOutputTokens(2048),
model.WithModelConfigToolChoice(true),
),
},
{
Model: "ERNIE-4.0-Turbo-8K",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaidu,
InputPrice: 0.02,
OutputPrice: 0.06,
RPM: 10000,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(6144),
model.WithModelConfigMaxInputTokens(6144),
model.WithModelConfigMaxOutputTokens(2048),
model.WithModelConfigToolChoice(true),
),
},
{
Model: "ERNIE-4.0-Turbo-128K",
@@ -84,24 +98,27 @@ var ModelList = []*model.ModelConfig{
Owner: model.ModelOwnerBaidu,
InputPrice: 0.02,
OutputPrice: 0.06,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 126976,
model.ModelConfigMaxInputTokensKey: 126976,
model.ModelConfigMaxOutputTokensKey: 4096,
},
RPM: 60,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(126976),
model.WithModelConfigMaxInputTokens(126976),
model.WithModelConfigMaxOutputTokens(4096),
model.WithModelConfigToolChoice(true),
),
},
{
Model: "ERNIE-3.5-8K-Preview",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaidu,
InputPrice: 0.0008,
OutputPrice: 0.002,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 5120,
model.ModelConfigMaxInputTokensKey: 5120,
model.ModelConfigMaxOutputTokensKey: 2048,
},
RPM: 300,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(5120),
model.WithModelConfigMaxInputTokens(5120),
model.WithModelConfigMaxOutputTokens(2048),
model.WithModelConfigToolChoice(true),
),
},
{
Model: "ERNIE-3.5-8K",
@@ -109,11 +126,13 @@ var ModelList = []*model.ModelConfig{
Owner: model.ModelOwnerBaidu,
InputPrice: 0.0008,
OutputPrice: 0.002,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 5120,
model.ModelConfigMaxInputTokensKey: 5120,
model.ModelConfigMaxOutputTokensKey: 2048,
},
RPM: 10000,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(5120),
model.WithModelConfigMaxInputTokens(5120),
model.WithModelConfigMaxOutputTokens(2048),
model.WithModelConfigToolChoice(true),
),
},
{
Model: "ERNIE-3.5-128K",
@@ -121,24 +140,26 @@ var ModelList = []*model.ModelConfig{
Owner: model.ModelOwnerBaidu,
InputPrice: 0.0008,
OutputPrice: 0.002,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 126976,
model.ModelConfigMaxInputTokensKey: 126976,
model.ModelConfigMaxOutputTokensKey: 4096,
},
RPM: 5000,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(126976),
model.WithModelConfigMaxInputTokens(126976),
model.WithModelConfigMaxOutputTokens(4096),
model.WithModelConfigToolChoice(true),
),
},
{
Model: "ERNIE-Speed-8K",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaidu,
InputPrice: 0.0001,
OutputPrice: 0.0001,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 7168,
model.ModelConfigMaxInputTokensKey: 7168,
model.ModelConfigMaxOutputTokensKey: 2048,
},
RPM: 500,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(7168),
model.WithModelConfigMaxInputTokens(7168),
model.WithModelConfigMaxOutputTokens(2048),
),
},
{
Model: "ERNIE-Speed-128K",
@@ -146,11 +167,12 @@ var ModelList = []*model.ModelConfig{
Owner: model.ModelOwnerBaidu,
InputPrice: 0.0001,
OutputPrice: 0.0001,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 126976,
model.ModelConfigMaxInputTokensKey: 126976,
model.ModelConfigMaxOutputTokensKey: 4096,
},
RPM: 500,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(126976),
model.WithModelConfigMaxInputTokens(126976),
model.WithModelConfigMaxOutputTokens(4096),
),
},
{
Model: "ERNIE-Speed-Pro-128K",
@@ -158,24 +180,25 @@ var ModelList = []*model.ModelConfig{
Owner: model.ModelOwnerBaidu,
InputPrice: 0.0003,
OutputPrice: 0.0006,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 126976,
model.ModelConfigMaxInputTokensKey: 126976,
model.ModelConfigMaxOutputTokensKey: 4096,
},
RPM: 10000,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(126976),
model.WithModelConfigMaxInputTokens(126976),
model.WithModelConfigMaxOutputTokens(4096),
),
},
{
Model: "ERNIE-Lite-8K",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaidu,
InputPrice: 0.0001,
OutputPrice: 0.0001,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 6144,
model.ModelConfigMaxInputTokensKey: 6144,
model.ModelConfigMaxOutputTokensKey: 2048,
},
RPM: 500,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(6144),
model.WithModelConfigMaxInputTokens(6144),
model.WithModelConfigMaxOutputTokens(2048),
),
},
{
Model: "ERNIE-Lite-Pro-128K",
@@ -183,37 +206,39 @@ var ModelList = []*model.ModelConfig{
Owner: model.ModelOwnerBaidu,
InputPrice: 0.0002,
OutputPrice: 0.0004,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 126976,
model.ModelConfigMaxInputTokensKey: 126976,
model.ModelConfigMaxOutputTokensKey: 4096,
},
RPM: 10000,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(126976),
model.WithModelConfigMaxInputTokens(126976),
model.WithModelConfigMaxOutputTokens(4096),
model.WithModelConfigToolChoice(true),
),
},
{
Model: "ERNIE-Tiny-8K",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaidu,
InputPrice: 0.0001,
OutputPrice: 0.0001,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 6144,
model.ModelConfigMaxInputTokensKey: 6144,
model.ModelConfigMaxOutputTokensKey: 2048,
},
RPM: 10000,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(6144),
model.WithModelConfigMaxInputTokens(6144),
model.WithModelConfigMaxOutputTokens(2048),
),
},
{
Model: "ERNIE-Character-8K",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaidu,
InputPrice: 0.0003,
OutputPrice: 0.0006,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 6144,
model.ModelConfigMaxInputTokensKey: 6144,
model.ModelConfigMaxOutputTokensKey: 2048,
},
RPM: 60,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(6144),
model.WithModelConfigMaxInputTokens(6144),
model.WithModelConfigMaxOutputTokens(2048),
),
},
{
Model: "ERNIE-Character-Fiction-8K",
@@ -221,23 +246,49 @@ var ModelList = []*model.ModelConfig{
Owner: model.ModelOwnerBaidu,
InputPrice: 0.0003,
OutputPrice: 0.0006,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 5120,
model.ModelConfigMaxInputTokensKey: 5120,
model.ModelConfigMaxOutputTokensKey: 2048,
},
RPM: 300,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(5120),
model.WithModelConfigMaxInputTokens(5120),
model.WithModelConfigMaxOutputTokens(2048),
),
},
{
Model: "ERNIE-Novel-8K",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerBaidu,
InputPrice: 0.04,
OutputPrice: 0.12,
Config: map[model.ModelConfigKey]any{
model.ModelConfigMaxContextTokensKey: 6144,
model.ModelConfigMaxInputTokensKey: 6144,
model.ModelConfigMaxOutputTokensKey: 2048,
},
RPM: 60,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(6144),
model.WithModelConfigMaxInputTokens(6144),
model.WithModelConfigMaxOutputTokens(2048),
),
},
{
Model: "DeepSeek-V3",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerDeepSeek,
InputPrice: 0.0008,
OutputPrice: 0.0016,
RPM: 1000,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(64000),
model.WithModelConfigMaxOutputTokens(8192),
),
},
{
Model: "DeepSeek-R1",
Type: relaymode.ChatCompletions,
Owner: model.ModelOwnerDeepSeek,
InputPrice: 0.002,
OutputPrice: 0.008,
RPM: 1000,
Config: model.NewModelConfig(
model.WithModelConfigMaxContextTokens(64000),
model.WithModelConfigMaxOutputTokens(8192),
),
},
}
@@ -16,6 +16,10 @@ type Adaptor struct {
const baseURL = "https://api.cloudflare.com"
func (a *Adaptor) GetBaseURL() string {
return baseURL
}
// WorkerAI cannot be used across accounts with AIGateWay
// https://developers.cloudflare.com/ai-gateway/providers/workersai/#openai-compatible-endpoints
// https://gateway.ai.cloudflare.com/v1/{account_id}/{gateway_id}/workers-ai
@@ -25,15 +29,12 @@ func isAIGateWay(baseURL string) bool {
func (a *Adaptor) GetRequestURL(meta *meta.Meta) (string, error) {
u := meta.Channel.BaseURL
if u == "" {
u = baseURL
}
isAIGateWay := isAIGateWay(u)
var urlPrefix string
if isAIGateWay {
urlPrefix = u
} else {
urlPrefix = fmt.Sprintf("%s/client/v4/accounts/%s/ai", u, meta.Channel.Config.UserID)
urlPrefix = fmt.Sprintf("%s/client/v4/accounts/%s/ai", u, meta.Channel.Key)
}
switch meta.Mode {
@@ -43,9 +44,9 @@ func (a *Adaptor) GetRequestURL(meta *meta.Meta) (string, error) {
return urlPrefix + "/v1/embeddings", nil
default:
if isAIGateWay {
return fmt.Sprintf("%s/%s", urlPrefix, meta.ActualModelName), nil
return fmt.Sprintf("%s/%s", urlPrefix, meta.ActualModel), nil
}
return fmt.Sprintf("%s/run/%s", urlPrefix, meta.ActualModelName), nil
return fmt.Sprintf("%s/run/%s", urlPrefix, meta.ActualModel), nil
}
}
+12 -12
View File
@@ -20,12 +20,12 @@ type Adaptor struct{}
const baseURL = "https://api.cohere.ai"
func (a *Adaptor) GetBaseURL() string {
return baseURL
}
func (a *Adaptor) GetRequestURL(meta *meta.Meta) (string, error) {
u := meta.Channel.BaseURL
if u == "" {
u = baseURL
}
return u + "/v1/chat", nil
return meta.Channel.BaseURL + "/v1/chat", nil
}
func (a *Adaptor) SetupRequestHeader(meta *meta.Meta, _ *gin.Context, req *http.Request) error {
@@ -33,21 +33,21 @@ func (a *Adaptor) SetupRequestHeader(meta *meta.Meta, _ *gin.Context, req *http.
return nil
}
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Reader, error) {
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (string, http.Header, io.Reader, error) {
request, err := utils.UnmarshalGeneralOpenAIRequest(req)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
request.Model = meta.ActualModelName
request.Model = meta.ActualModel
requestBody := ConvertRequest(request)
if requestBody == nil {
return nil, nil, errors.New("request body is nil")
return "", nil, nil, errors.New("request body is nil")
}
data, err := json.Marshal(requestBody)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
return nil, bytes.NewReader(data), nil
return http.MethodPost, nil, bytes.NewReader(data), nil
}
func (a *Adaptor) DoRequest(_ *meta.Meta, _ *gin.Context, req *http.Request) (*http.Response, error) {
@@ -62,7 +62,7 @@ func (a *Adaptor) DoResponse(meta *meta.Meta, c *gin.Context, resp *http.Respons
if utils.IsStreamResponse(resp) {
err, usage = StreamHandler(c, resp)
} else {
err, usage = Handler(c, resp, meta.PromptTokens, meta.ActualModelName)
err, usage = Handler(c, resp, meta.InputTokens, meta.ActualModel)
}
}
return
+8 -11
View File
@@ -5,16 +5,16 @@ import (
"fmt"
"net/http"
"strings"
"time"
"github.com/gin-gonic/gin"
json "github.com/json-iterator/go"
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/common/conv"
"github.com/labring/sealos/service/aiproxy/common/render"
"github.com/labring/sealos/service/aiproxy/middleware"
"github.com/gin-gonic/gin"
"github.com/labring/sealos/service/aiproxy/common"
"github.com/labring/sealos/service/aiproxy/common/helper"
"github.com/labring/sealos/service/aiproxy/relay/adaptor/openai"
"github.com/labring/sealos/service/aiproxy/relay/constant"
"github.com/labring/sealos/service/aiproxy/relay/model"
)
@@ -26,7 +26,7 @@ func stopReasonCohere2OpenAI(reason *string) string {
}
switch *reason {
case "COMPLETE":
return "stop"
return constant.StopFinishReason
default:
return *reason
}
@@ -125,7 +125,7 @@ func ResponseCohere2OpenAI(cohereResponse *Response) *openai.TextResponse {
ID: "chatcmpl-" + cohereResponse.ResponseID,
Model: "model",
Object: "chat.completion",
Created: helper.GetTimestamp(),
Created: time.Now().Unix(),
Choices: []*openai.TextResponseChoice{&choice},
}
return &fullTextResponse
@@ -136,7 +136,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC
log := middleware.GetLogger(c)
createdTime := helper.GetTimestamp()
createdTime := time.Now().Unix()
scanner := bufio.NewScanner(resp.Body)
scanner.Split(bufio.ScanLines)
@@ -168,10 +168,7 @@ func StreamHandler(c *gin.Context, resp *http.Response) (*model.ErrorWithStatusC
response.Model = c.GetString("original_model")
response.Created = createdTime
err = render.ObjectData(c, response)
if err != nil {
log.Error("error rendering stream response: " + err.Error())
}
_ = render.ObjectData(c, response)
}
if err := scanner.Err(); err != nil {
+44 -26
View File
@@ -5,6 +5,7 @@ import (
"errors"
"io"
"net/http"
"strings"
"github.com/gin-gonic/gin"
json "github.com/json-iterator/go"
@@ -12,6 +13,7 @@ import (
"github.com/labring/sealos/service/aiproxy/relay/adaptor/openai"
"github.com/labring/sealos/service/aiproxy/relay/meta"
relaymodel "github.com/labring/sealos/service/aiproxy/relay/model"
"github.com/labring/sealos/service/aiproxy/relay/relaymode"
"github.com/labring/sealos/service/aiproxy/relay/utils"
)
@@ -19,42 +21,58 @@ type Adaptor struct{}
const baseURL = "https://api.coze.com"
func (a *Adaptor) GetBaseURL() string {
return baseURL
}
func (a *Adaptor) GetRequestURL(meta *meta.Meta) (string, error) {
u := meta.Channel.BaseURL
if u == "" {
u = baseURL
}
return u + "/open_api/v2/chat", nil
return meta.Channel.BaseURL + "/open_api/v2/chat", nil
}
func (a *Adaptor) SetupRequestHeader(meta *meta.Meta, _ *gin.Context, req *http.Request) error {
req.Header.Set("Authorization", "Bearer "+meta.Channel.Key)
token, _, err := getTokenAndUserID(meta.Channel.Key)
if err != nil {
return err
}
req.Header.Set("Authorization", "Bearer "+token)
return nil
}
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (http.Header, io.Reader, error) {
func (a *Adaptor) ConvertRequest(meta *meta.Meta, req *http.Request) (string, http.Header, io.Reader, error) {
if meta.Mode != relaymode.ChatCompletions {
return "", nil, nil, errors.New("coze only support chat completions")
}
request, err := utils.UnmarshalGeneralOpenAIRequest(req)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
request.User = meta.Channel.Config.UserID
request.Model = meta.ActualModelName
requestBody := ConvertRequest(request)
if requestBody == nil {
return nil, nil, errors.New("request body is nil")
}
data, err := json.Marshal(requestBody)
_, userID, err := getTokenAndUserID(meta.Channel.Key)
if err != nil {
return nil, nil, err
return "", nil, nil, err
}
return nil, bytes.NewReader(data), nil
}
func (a *Adaptor) ConvertImageRequest(request *relaymodel.ImageRequest) (any, error) {
if request == nil {
return nil, errors.New("request is nil")
request.User = userID
request.Model = meta.ActualModel
cozeRequest := Request{
Stream: request.Stream,
User: request.User,
BotID: strings.TrimPrefix(meta.ActualModel, "bot-"),
}
return request, nil
for i, message := range request.Messages {
if i == len(request.Messages)-1 {
cozeRequest.Query = message.StringContent()
continue
}
cozeMessage := Message{
Role: message.Role,
Content: message.StringContent(),
}
cozeRequest.ChatHistory = append(cozeRequest.ChatHistory, cozeMessage)
}
data, err := json.Marshal(cozeRequest)
if err != nil {
return "", nil, nil, err
}
return http.MethodPost, nil, bytes.NewReader(data), nil
}
func (a *Adaptor) DoRequest(_ *meta.Meta, _ *gin.Context, req *http.Request) (*http.Response, error) {
@@ -66,14 +84,14 @@ func (a *Adaptor) DoResponse(meta *meta.Meta, c *gin.Context, resp *http.Respons
if utils.IsStreamResponse(resp) {
err, responseText = StreamHandler(c, resp)
} else {
err, responseText = Handler(c, resp, meta.PromptTokens, meta.ActualModelName)
err, responseText = Handler(c, resp, meta.InputTokens, meta.ActualModel)
}
if responseText != nil {
usage = openai.ResponseText2Usage(*responseText, meta.ActualModelName, meta.PromptTokens)
usage = openai.ResponseText2Usage(*responseText, meta.ActualModel, meta.InputTokens)
} else {
usage = &relaymodel.Usage{}
}
usage.PromptTokens = meta.PromptTokens
usage.PromptTokens = meta.InputTokens
usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens
return
}
+30
View File
@@ -0,0 +1,30 @@
package coze
import (
"errors"
"strings"
"github.com/labring/sealos/service/aiproxy/relay/adaptor"
)
var _ adaptor.KeyValidator = (*Adaptor)(nil)
func (a *Adaptor) ValidateKey(key string) error {
_, _, err := getTokenAndUserID(key)
if err != nil {
return err
}
return nil
}
func (a *Adaptor) KeyHelp() string {
return "token|user_id"
}
func getTokenAndUserID(key string) (string, string, error) {
split := strings.Split(key, "|")
if len(split) != 2 {
return "", "", errors.New("invalid key format")
}
return split[0], split[1], nil
}

Some files were not shown because too many files have changed in this diff Show More