Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions VERSION
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
v1.0.0-rc.41
10 changes: 10 additions & 0 deletions controller/log.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import (

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
"github.com/QuantumNous/new-api/service"

"github.com/gin-gonic/gin"
)
Expand All @@ -32,6 +33,7 @@ func GetAllLogs(c *gin.Context) {
} else {
model.FormatRootLogs(logs)
}
normalizeUsagePromptAuditLogs(logs)
pageInfo.SetTotal(int(total))
pageInfo.SetItems(logs)
common.ApiSuccess(c, pageInfo)
Expand All @@ -54,6 +56,7 @@ func GetUserLogs(c *gin.Context) {
common.ApiError(c, err)
return
}
normalizeUsagePromptAuditLogs(logs)
pageInfo.SetTotal(int(total))
pageInfo.SetItems(logs)
common.ApiSuccess(c, pageInfo)
Expand Down Expand Up @@ -93,13 +96,20 @@ func GetLogByKey(c *gin.Context) {
})
return
}
normalizeUsagePromptAuditLogs(logs)
c.JSON(200, gin.H{
"success": true,
"message": "",
"data": logs,
})
}

func normalizeUsagePromptAuditLogs(logs []*model.Log) {
for _, log := range logs {
log.Other = service.NormalizeUsagePromptAuditOther(log.Other)
}
}

func GetLogsStat(c *gin.Context) {
logType, _ := strconv.Atoi(c.Query("type"))
startTimestamp, _ := strconv.ParseInt(c.Query("start_timestamp"), 10, 64)
Expand Down
11 changes: 11 additions & 0 deletions controller/relay.go
Original file line number Diff line number Diff line change
Expand Up @@ -123,6 +123,17 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) {
newAPIError = types.NewError(err, types.ErrorCodeGenRelayInfoFailed)
return
}
if bodyStorage, bodyErr := common.GetBodyStorage(c); bodyErr == nil {
body, truncated, captureErr := service.CaptureUsagePromptAuditBodyFromStorage(bodyStorage)
if captureErr != nil {
logger.LogWarn(c, "failed to capture usage prompt audit client request body: "+captureErr.Error())
} else {
relayInfo.UsageLogRawClientRequestBody = body
relayInfo.UsageLogRawClientRequestBodyTruncated = truncated
}
} else {
logger.LogWarn(c, "failed to get request body storage for usage prompt audit: "+bodyErr.Error())
}

defer func() {
recovered := recover()
Expand Down
7 changes: 4 additions & 3 deletions controller/token.go
Original file line number Diff line number Diff line change
Expand Up @@ -129,13 +129,13 @@ func setTokenAutoGroups(c *gin.Context, token *model.Token, groups []string) boo

func GetAllTokens(c *gin.Context) {
userId := c.GetInt("id")
groups := c.QueryArray("group")
pageInfo := common.GetPageQuery(c)
tokens, err := model.GetAllUserTokens(userId, pageInfo.GetStartIdx(), pageInfo.GetPageSize())
tokens, total, err := model.GetAllUserTokens(userId, groups, pageInfo.GetStartIdx(), pageInfo.GetPageSize())
if err != nil {
common.ApiError(c, err)
return
}
total, _ := model.CountUserTokens(userId)
pageInfo.SetTotal(int(total))
pageInfo.SetItems(buildMaskedTokenResponses(tokens))
common.ApiSuccess(c, pageInfo)
Expand All @@ -145,10 +145,11 @@ func SearchTokens(c *gin.Context) {
userId := c.GetInt("id")
keyword := c.Query("keyword")
token := c.Query("token")
groups := c.QueryArray("group")

pageInfo := common.GetPageQuery(c)

tokens, total, err := model.SearchUserTokens(userId, keyword, token, pageInfo.GetStartIdx(), pageInfo.GetPageSize())
tokens, total, err := model.SearchUserTokens(userId, keyword, token, groups, pageInfo.GetStartIdx(), pageInfo.GetPageSize())
if err != nil {
common.ApiError(c, err)
return
Expand Down
129 changes: 129 additions & 0 deletions controller/token_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@ type tokenAPIResponse struct {

type tokenPageResponse struct {
Items []tokenResponseItem `json:"items"`
Total int `json:"total"`
}

type tokenResponseItem struct {
Expand Down Expand Up @@ -84,6 +85,10 @@ func openTokenControllerTestDB(t *testing.T) *gorm.DB {
common.SetDatabaseTypes(common.DatabaseTypeSQLite, common.DatabaseTypeSQLite)
common.RedisEnabled = false

// model 的 commonGroupCol/commonKeyCol 等方言列名只在 InitDB 里初始化,
// 隔离运行(go test -run)时为空会导致生成的 SQL 缺列名。
initModelListColumnNames(t)

dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", strings.ReplaceAll(t.Name(), "/", "_"))
db, err := gorm.Open(sqlite.Open(dsn), &gorm.Config{})
if err != nil {
Expand Down Expand Up @@ -490,6 +495,130 @@ func TestSearchTokensMasksKeyInResponse(t *testing.T) {
}
}

func TestSearchTokensKeywordMatchesNameSubstring(t *testing.T) {
db := setupTokenControllerTestDB(t)
seedToken(t, db, 1, "north-beijing-token", "abcd1234efgh5678")
seedToken(t, db, 1, "shanghai-token", "mnop1234qrst5678")

ctx, recorder := newAuthenticatedContext(t, http.MethodGet, "/api/token/search?keyword=beijing&p=1&size=10", nil, 1)
SearchTokens(ctx)

response := decodeAPIResponse(t, recorder)
require.True(t, response.Success, "expected success response, got message: %s", response.Message)

var page tokenPageResponse
require.NoError(t, common.Unmarshal(response.Data, &page))
require.Len(t, page.Items, 1)
require.Equal(t, "north-beijing-token", page.Items[0].Name)
}

func TestSearchTokensTokenMatchesKeySubstring(t *testing.T) {
db := setupTokenControllerTestDB(t)
seedToken(t, db, 1, "matching-key-token", "abcd1234efgh5678")
seedToken(t, db, 1, "other-key-token", "mnop1234qrst5678")

ctx, recorder := newAuthenticatedContext(t, http.MethodGet, "/api/token/search?token=efgh&p=1&size=10", nil, 1)
SearchTokens(ctx)

response := decodeAPIResponse(t, recorder)
require.True(t, response.Success, "expected success response, got message: %s", response.Message)

var page tokenPageResponse
require.NoError(t, common.Unmarshal(response.Data, &page))
require.Len(t, page.Items, 1)
require.Equal(t, "matching-key-token", page.Items[0].Name)
}

func TestGetAllTokensFiltersByGroup(t *testing.T) {
db := setupTokenControllerTestDB(t)
premium := seedToken(t, db, 1, "premium-token", "abcd1234efgh5678")
require.NoError(t, db.Model(premium).Update("group", "premium").Error)
seedToken(t, db, 1, "default-token", "mnop1234qrst5678")
otherUser := seedToken(t, db, 2, "other-user-premium", "uvwx1234yzab5678")
require.NoError(t, db.Model(otherUser).Update("group", "premium").Error)

ctx, recorder := newAuthenticatedContext(t, http.MethodGet, "/api/token/?group=premium&p=1&size=10", nil, 1)
GetAllTokens(ctx)

response := decodeAPIResponse(t, recorder)
require.True(t, response.Success, "expected success response, got message: %s", response.Message)

var page tokenPageResponse
require.NoError(t, common.Unmarshal(response.Data, &page))
require.Equal(t, 1, page.Total)
require.Len(t, page.Items, 1)
require.Equal(t, "premium-token", page.Items[0].Name)
}

func TestGetAllTokensFiltersByMultipleGroups(t *testing.T) {
db := setupTokenControllerTestDB(t)
premium := seedToken(t, db, 1, "premium-token", "abcd1234efgh5678")
require.NoError(t, db.Model(premium).Update("group", "premium").Error)
standard := seedToken(t, db, 1, "standard-token", "mnop1234qrst5678")
require.NoError(t, db.Model(standard).Update("group", "standard").Error)
seedToken(t, db, 1, "default-token", "ijkl1234mnop5678")
otherUser := seedToken(t, db, 2, "other-user-premium", "uvwx1234yzab5678")
require.NoError(t, db.Model(otherUser).Update("group", "premium").Error)

ctx, recorder := newAuthenticatedContext(t, http.MethodGet, "/api/token/?group=premium&group=standard&p=1&size=10", nil, 1)
GetAllTokens(ctx)

response := decodeAPIResponse(t, recorder)
require.True(t, response.Success, "expected success response, got message: %s", response.Message)

var page tokenPageResponse
require.NoError(t, common.Unmarshal(response.Data, &page))
require.Equal(t, 2, page.Total)
require.Len(t, page.Items, 2)
assert.ElementsMatch(t, []string{"premium-token", "standard-token"}, []string{
page.Items[0].Name,
page.Items[1].Name,
})
}

func TestSearchTokensCombinesGroupAndNameSubstring(t *testing.T) {
db := setupTokenControllerTestDB(t)
premium := seedToken(t, db, 1, "north-beijing-premium", "abcd1234efgh5678")
require.NoError(t, db.Model(premium).Update("group", "premium").Error)
seedToken(t, db, 1, "north-beijing-default", "mnop1234qrst5678")

ctx, recorder := newAuthenticatedContext(t, http.MethodGet, "/api/token/search?keyword=beijing&group=premium&p=1&size=10", nil, 1)
SearchTokens(ctx)

response := decodeAPIResponse(t, recorder)
require.True(t, response.Success, "expected success response, got message: %s", response.Message)

var page tokenPageResponse
require.NoError(t, common.Unmarshal(response.Data, &page))
require.Equal(t, 1, page.Total)
require.Len(t, page.Items, 1)
require.Equal(t, "north-beijing-premium", page.Items[0].Name)
}

func TestSearchTokensCombinesMultipleGroupsAndNameSubstring(t *testing.T) {
db := setupTokenControllerTestDB(t)
premium := seedToken(t, db, 1, "north-beijing-premium", "abcd1234efgh5678")
require.NoError(t, db.Model(premium).Update("group", "premium").Error)
standard := seedToken(t, db, 1, "north-beijing-standard", "mnop1234qrst5678")
require.NoError(t, db.Model(standard).Update("group", "standard").Error)
seedToken(t, db, 1, "north-beijing-default", "ijkl1234mnop5678")

ctx, recorder := newAuthenticatedContext(t, http.MethodGet, "/api/token/search?keyword=beijing&group=premium&group=standard&p=1&size=10", nil, 1)
SearchTokens(ctx)

response := decodeAPIResponse(t, recorder)
require.True(t, response.Success, "expected success response, got message: %s", response.Message)

var page tokenPageResponse
require.NoError(t, common.Unmarshal(response.Data, &page))
require.Equal(t, 2, page.Total)
require.Len(t, page.Items, 2)
assert.ElementsMatch(t, []string{"north-beijing-premium", "north-beijing-standard"}, []string{
page.Items[0].Name,
page.Items[1].Name,
})
}

func TestGetTokenMasksKeyInResponse(t *testing.T) {
db := setupTokenControllerTestDB(t)
token := seedToken(t, db, 1, "detail-token", "qrst1234uvwx5678")
Expand Down
2 changes: 1 addition & 1 deletion model/ability.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@ import (
)

type Ability struct {
Group string `json:"group" gorm:"type:varchar(64);primaryKey;autoIncrement:false"`
Group string `json:"group" gorm:"type:varchar(255);primaryKey;autoIncrement:false"`
Model string `json:"model" gorm:"type:varchar(255);primaryKey;autoIncrement:false"`
ChannelId int `json:"channel_id" gorm:"primaryKey;autoIncrement:false;index"`
Enabled bool `json:"enabled"`
Expand Down
2 changes: 1 addition & 1 deletion model/channel.go
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@ type Channel struct {
Balance float64 `json:"balance"` // in USD
BalanceUpdatedTime int64 `json:"balance_updated_time" gorm:"bigint"`
Models string `json:"models"`
Group string `json:"group" gorm:"type:varchar(64);default:'default'"`
Group string `json:"group" gorm:"type:varchar(255);default:'default'"`
UsedQuota int64 `json:"used_quota" gorm:"bigint;default:0"`
ModelMapping *string `json:"model_mapping" gorm:"type:text"`
//MaxInputTokens *int `json:"max_input_tokens" gorm:"default:0"`
Expand Down
39 changes: 29 additions & 10 deletions model/token.go
Original file line number Diff line number Diff line change
Expand Up @@ -103,11 +103,16 @@ func (token *Token) GetIpLimits() []string {
return ipLimits
}

func GetAllUserTokens(userId int, startIdx int, num int) ([]*Token, error) {
var tokens []*Token
var err error
err = DB.Where("user_id = ?", userId).Order("id desc").Limit(num).Offset(startIdx).Find(&tokens).Error
return tokens, err
func GetAllUserTokens(userId int, groups []string, startIdx int, num int) (tokens []*Token, total int64, err error) {
query := DB.Model(&Token{}).Where("user_id = ?", userId)
if len(groups) > 0 {
query = query.Where(commonGroupCol+" IN ?", groups)
}
if err = query.Count(&total).Error; err != nil {
return nil, 0, err
}
err = query.Order("id desc").Limit(num).Offset(startIdx).Find(&tokens).Error
return tokens, total, err
}

// sanitizeLikePattern 校验并清洗用户输入的 LIKE 搜索模式。
Expand Down Expand Up @@ -154,9 +159,20 @@ func validateLikePattern(input string) error {
return nil
}

func sanitizeContainsLikePattern(input string) (string, error) {
pattern, err := sanitizeLikePattern(input)
if err != nil {
return "", err
}
if strings.Contains(pattern, "%") {
return pattern, nil
}
return "%" + pattern + "%", nil
}

const searchHardLimit = 100

func SearchUserTokens(userId int, keyword string, token string, offset int, limit int) (tokens []*Token, total int64, err error) {
func SearchUserTokens(userId int, keyword string, token string, groups []string, offset int, limit int) (tokens []*Token, total int64, err error) {
// model 层强制截断
if limit <= 0 || limit > searchHardLimit {
limit = searchHardLimit
Expand All @@ -171,30 +187,33 @@ func SearchUserTokens(userId int, keyword string, token string, offset int, limi

// 超量用户(令牌数超过上限)只允许精确搜索,禁止模糊搜索
maxTokens := operation_setting.GetMaxUserTokens()
hasFuzzy := strings.Contains(keyword, "%") || strings.Contains(token, "%")
hasFuzzy := keyword != "" || token != ""
if hasFuzzy {
count, err := CountUserTokens(userId)
if err != nil {
common.SysLog("failed to count user tokens: " + err.Error())
return nil, 0, errors.New("获取令牌数量失败")
}
if int(count) > maxTokens {
return nil, 0, errors.New("令牌数量超过上限,仅允许精确搜索,请勿使用 % 通配符")
return nil, 0, errors.New("令牌数量超过上限,仅允许精确搜索")
}
}

baseQuery := DB.Model(&Token{}).Where("user_id = ?", userId)
if len(groups) > 0 {
baseQuery = baseQuery.Where(commonGroupCol+" IN ?", groups)
}

// 非空才加 LIKE 条件,空则跳过(不过滤该字段)
if keyword != "" {
keywordPattern, err := sanitizeLikePattern(keyword)
keywordPattern, err := sanitizeContainsLikePattern(keyword)
if err != nil {
return nil, 0, err
}
baseQuery = baseQuery.Where("name LIKE ? ESCAPE '!'", keywordPattern)
}
if token != "" {
tokenPattern, err := sanitizeLikePattern(token)
tokenPattern, err := sanitizeContainsLikePattern(token)
if err != nil {
return nil, 0, err
}
Expand Down
2 changes: 1 addition & 1 deletion model/user.go
Original file line number Diff line number Diff line change
Expand Up @@ -97,7 +97,7 @@ type User struct {
Quota int `json:"quota" gorm:"type:int;default:0"`
UsedQuota int `json:"used_quota" gorm:"type:int;default:0;column:used_quota"` // used quota
RequestCount int `json:"request_count" gorm:"type:int;default:0;"` // request number
Group string `json:"group" gorm:"type:varchar(64);default:'default'"`
Group string `json:"group" gorm:"type:varchar(255);default:'default'"`
AffCode string `json:"aff_code" gorm:"type:varchar(32);column:aff_code;uniqueIndex"`
AffCount int `json:"aff_count" gorm:"type:int;default:0;column:aff_count"`
AffQuota int `json:"aff_quota" gorm:"type:int;default:0;column:aff_quota"` // 邀请剩余额度
Expand Down
13 changes: 13 additions & 0 deletions relay/channel/claude/usage_prompt_audit.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,13 @@
package claude

import (
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/service"
)

func recordUsagePromptAuditResponseBody(info *relaycommon.RelayInfo, body []byte) {
if info == nil || len(body) == 0 {
return
}
info.UsageLogRawResponseBody, info.UsageLogRawResponseBodyTruncated = service.CaptureUsagePromptAuditBytes(body)
}
2 changes: 2 additions & 0 deletions relay/channel/openai/relay-openai.go
Original file line number Diff line number Diff line change
Expand Up @@ -186,6 +186,7 @@ func OaiStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Re
}

applyUsagePostProcessing(info, usage, common.StringToByteSlice(usageFrame))
recordUsageLogSyntheticChatStreamResponse(info, responseTextBuilder.String(), lastStreamData)

for _, name := range streamFunctionCallNames {
info.CountBillableToolCall(dto.BuildInCallFunctionCall, name)
Expand Down Expand Up @@ -272,6 +273,7 @@ func OpenaiHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Respo
return nil, types.NewOpenAIError(fmt.Errorf("openrouter response success=false"), types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
}
}
recordUsageLogRawResponseBody(info, responseBody)

err = common.Unmarshal(responseBody, &simpleResponse)
if err != nil {
Expand Down
Loading
Loading