mirror of
https://github.com/labring/sealos.git
synced 2026-09-24 15:46:19 +08:00
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:
@@ -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"]
|
||||
|
||||
@@ -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`
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
)
|
||||
|
||||
Vendored
+52
-9
@@ -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"
|
||||
)
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -1,7 +0,0 @@
|
||||
package helper
|
||||
|
||||
type Key string
|
||||
|
||||
const (
|
||||
RequestIDKey Key = "X-Request-Id"
|
||||
)
|
||||
@@ -1,9 +0,0 @@
|
||||
package helper
|
||||
|
||||
import (
|
||||
"time"
|
||||
)
|
||||
|
||||
func GetTimestamp() int64 {
|
||||
return time.Now().Unix()
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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]
|
||||
}
|
||||
@@ -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})
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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,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 {
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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"`
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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")
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
+884
-354
File diff suppressed because it is too large
Load Diff
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
@@ -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())
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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),
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
Reference in New Issue
Block a user