Skip to content
Open
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
17 changes: 17 additions & 0 deletions controller/token.go
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,23 @@ func SearchTokens(c *gin.Context) {
common.ApiSuccess(c, pageInfo)
}

// SearchAllTokens 管理员全局搜索所有用户的 API KEY。
func SearchAllTokens(c *gin.Context) {
keyword := c.Query("keyword")
token := c.Query("token")

pageInfo := common.GetPageQuery(c)

tokens, total, err := model.SearchAllTokens(keyword, token, pageInfo.GetStartIdx(), pageInfo.GetPageSize())
if err != nil {
common.ApiError(c, err)
return
}
pageInfo.SetTotal(int(total))
pageInfo.SetItems(buildMaskedTokenResponses(tokens))
common.ApiSuccess(c, pageInfo)
}

func GetToken(c *gin.Context) {
id, err := strconv.Atoi(c.Param("id"))
userId := c.GetInt("id")
Expand Down
6 changes: 4 additions & 2 deletions controller/usedata.go
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,8 @@ func GetAllQuotaDates(c *gin.Context) {
startTimestamp, _ := strconv.ParseInt(c.Query("start_timestamp"), 10, 64)
endTimestamp, _ := strconv.ParseInt(c.Query("end_timestamp"), 10, 64)
username := c.Query("username")
dates, err := model.GetAllQuotaDates(startTimestamp, endTimestamp, username)
tokenID, _ := strconv.Atoi(c.Query("token_id"))
dates, err := model.GetAllQuotaDates(startTimestamp, endTimestamp, username, tokenID)
if err != nil {
common.ApiError(c, err)
return
Expand Down Expand Up @@ -64,6 +65,7 @@ func GetUserQuotaDates(c *gin.Context) {
userId := c.GetInt("id")
startTimestamp, _ := strconv.ParseInt(c.Query("start_timestamp"), 10, 64)
endTimestamp, _ := strconv.ParseInt(c.Query("end_timestamp"), 10, 64)
tokenID, _ := strconv.Atoi(c.Query("token_id"))
// 判断时间跨度是否超过 1 个月
if endTimestamp-startTimestamp > 2592000 {
c.JSON(http.StatusOK, gin.H{
Expand All @@ -72,7 +74,7 @@ func GetUserQuotaDates(c *gin.Context) {
})
return
}
dates, err := model.GetQuotaDataByUserId(userId, startTimestamp, endTimestamp)
dates, err := model.GetQuotaDataByUserId(userId, startTimestamp, endTimestamp, tokenID)
if err != nil {
common.ApiError(c, err)
return
Expand Down
87 changes: 87 additions & 0 deletions controller/usedata_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,87 @@
package controller

import (
"net/http"
"net/http/httptest"
"testing"

"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/model"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)

type quotaDatesResponse struct {
Success bool `json:"success"`
Message string `json:"message"`
Data []model.QuotaData `json:"data"`
}

func decodeQuotaDatesResponse(t *testing.T, recorder *httptest.ResponseRecorder) quotaDatesResponse {
t.Helper()
require.Equal(t, http.StatusOK, recorder.Code)
var payload quotaDatesResponse
require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &payload))
require.True(t, payload.Success, payload.Message)
return payload
}

func TestGetAllQuotaDatesFiltersByTokenID(t *testing.T) {
setupFlowControllerTestDB(t)

recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
ctx.Request = httptest.NewRequest(http.MethodGet, "/api/data?start_timestamp=1000&end_timestamp=2000&token_id=11", nil)

GetAllQuotaDates(ctx)

payload := decodeQuotaDatesResponse(t, recorder)
require.Len(t, payload.Data, 1)
require.Equal(t, "gpt-a", payload.Data[0].ModelName)
require.Equal(t, 2, payload.Data[0].Count)
require.Equal(t, 100, payload.Data[0].Quota)
}

func TestGetAllQuotaDatesIgnoresZeroTokenID(t *testing.T) {
setupFlowControllerTestDB(t)

recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
ctx.Request = httptest.NewRequest(http.MethodGet, "/api/data?start_timestamp=1000&end_timestamp=2000&token_id=0", nil)

GetAllQuotaDates(ctx)

payload := decodeQuotaDatesResponse(t, recorder)
require.Len(t, payload.Data, 2)
}

func TestGetUserQuotaDatesFiltersByTokenID(t *testing.T) {
setupFlowControllerTestDB(t)

recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
ctx.Set("id", 1)
ctx.Request = httptest.NewRequest(http.MethodGet, "/api/data/self?start_timestamp=1000&end_timestamp=2000&token_id=11", nil)

GetUserQuotaDates(ctx)

payload := decodeQuotaDatesResponse(t, recorder)
require.Len(t, payload.Data, 1)
require.Equal(t, "gpt-a", payload.Data[0].ModelName)
require.Equal(t, "alice", payload.Data[0].Username)
}

func TestGetUserQuotaDatesIgnoresOtherUserTokenID(t *testing.T) {
setupFlowControllerTestDB(t)

recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
ctx.Set("id", 1)
// token_id=22 属于 user 2,当前用户是 user 1,因此应返回空
ctx.Request = httptest.NewRequest(http.MethodGet, "/api/data/self?start_timestamp=1000&end_timestamp=2000&token_id=22", nil)

GetUserQuotaDates(ctx)

payload := decodeQuotaDatesResponse(t, recorder)
require.Empty(t, payload.Data)
}
79 changes: 68 additions & 11 deletions model/token.go
Original file line number Diff line number Diff line change
Expand Up @@ -118,14 +118,6 @@ func validateLikePattern(input string) error {
return errors.New("搜索模式中最多允许包含 2 个 % 通配符")
}

// 3. 含 % 时,去掉 % 后关键词长度必须 >= 2
if count > 0 {
stripped := strings.ReplaceAll(input, "%", "")
if len(stripped) < 2 {
return errors.New("使用模糊搜索时,关键词长度至少为 2 个字符")
}
}

return nil
}

Expand Down Expand Up @@ -160,20 +152,25 @@ func SearchUserTokens(userId int, keyword string, token string, offset int, limi

baseQuery := DB.Model(&Token{}).Where("user_id = ?", userId)

// 非空才加 LIKE 条件,空则跳过(不过滤该字段)
// 非空才加 LIKE 条件,空则跳过(不过滤该字段)。
// 若用户未显式输入通配符 %,默认按前缀模糊匹配(例如 RD11 匹配 RD1141)。
// 使用 LOWER() 实现跨数据库的大小写不敏感搜索。
if keyword != "" {
if !strings.Contains(keyword, "%") {
keyword = keyword + "%"
}
keywordPattern, err := sanitizeLikePattern(keyword)
if err != nil {
return nil, 0, err
}
baseQuery = baseQuery.Where("name LIKE ? ESCAPE '!'", keywordPattern)
baseQuery = baseQuery.Where("LOWER(name) LIKE LOWER(?) ESCAPE '!'", keywordPattern)
}
if token != "" {
tokenPattern, err := sanitizeLikePattern(token)
if err != nil {
return nil, 0, err
}
baseQuery = baseQuery.Where(commonKeyCol+" LIKE ? ESCAPE '!'", tokenPattern)
baseQuery = baseQuery.Where("LOWER("+commonKeyCol+") LIKE LOWER(?) ESCAPE '!'", tokenPattern)
}

// 先查匹配总数(用于分页,受 maxTokens 上限保护,避免全表 COUNT)
Expand All @@ -192,6 +189,66 @@ func SearchUserTokens(userId int, keyword string, token string, offset int, limi
return tokens, total, nil
}

// SearchAllTokens 全局搜索所有用户的 API KEY(管理员专用)。
// 实现复用 SearchUserTokens 的 LIKE 转义与截断逻辑,但不做 user_id 限制。
func SearchAllTokens(keyword string, token string, offset int, limit int) (tokens []*Token, total int64, err error) {
// model 层强制截断
if limit <= 0 || limit > searchHardLimit {
limit = searchHardLimit
}
if offset < 0 {
offset = 0
}

if token != "" {
token = strings.TrimPrefix(token, "sk-")
}

// 与 SearchUserTokens 一致:用 maxTokens 封顶 COUNT,避免 admin 全局搜索时
// 在大表上做全表 COUNT 扫描。total 实际为 min(真实匹配数, maxTokens),配合硬上限分页足够。
maxTokens := operation_setting.GetMaxUserTokens()

baseQuery := DB.Model(&Token{})

// 非空才加 LIKE 条件,空则跳过(不过滤该字段)。
// 若用户未显式输入通配符 %,默认按前缀模糊匹配(例如 RD11 匹配 RD1141)。
// 使用 LOWER() 实现跨数据库的大小写不敏感搜索。
if keyword != "" {
if !strings.Contains(keyword, "%") {
keyword = keyword + "%"
}
keywordPattern, err := sanitizeLikePattern(keyword)
if err != nil {
return nil, 0, err
}
baseQuery = baseQuery.Where("LOWER(name) LIKE LOWER(?) ESCAPE '!'", keywordPattern)
}
if token != "" {
tokenPattern, err := sanitizeLikePattern(token)
if err != nil {
return nil, 0, err
}
baseQuery = baseQuery.Where("LOWER("+commonKeyCol+") LIKE LOWER(?) ESCAPE '!'", tokenPattern)
}

// 先查匹配总数
// 与 SearchUserTokens 保持一致:用 maxTokens 封顶 COUNT,避免 admin 全局搜索时
// 在大表上做全表 COUNT 扫描。total 实际为 min(真实匹配数, maxTokens),配合硬上限分页足够。
err = baseQuery.Limit(maxTokens).Count(&total).Error
if err != nil {
Comment thread
coderabbitai[bot] marked this conversation as resolved.
common.SysError("failed to count search all tokens: " + err.Error())
return nil, 0, errors.New("搜索令牌失败")
}

// 再分页查数据
err = baseQuery.Order("id desc").Offset(offset).Limit(limit).Find(&tokens).Error
if err != nil {
common.SysError("failed to search all tokens: " + err.Error())
return nil, 0, errors.New("搜索令牌失败")
}
return tokens, total, nil
}

func ValidateUserToken(key string) (token *Token, err error) {
if key == "" {
return nil, ErrTokenNotProvided
Expand Down
39 changes: 26 additions & 13 deletions model/usedata.go
Original file line number Diff line number Diff line change
Expand Up @@ -138,24 +138,32 @@ func increaseQuotaData(quotaData *QuotaData) {
}
}

func GetQuotaDataByUsername(username string, startTime int64, endTime int64) (quotaData []*QuotaData, err error) {
// GetQuotaDataByUsername 根据用户名查询配额数据;传入 tokenId > 0 时进一步按 API KEY 过滤。
func GetQuotaDataByUsername(username string, startTime int64, endTime int64, tokenId int) (quotaData []*QuotaData, err error) {
var quotaDatas []*QuotaData
// 从quota_data表中查询数据
err = DB.Table("quota_data").
query := DB.Table("quota_data").
Select("user_id, username, model_name, created_at, sum(count) as count, sum(quota) as quota, sum(token_used) as token_used").
Where("username = ? and created_at >= ? and created_at <= ?", username, startTime, endTime).
Group("user_id, username, model_name, created_at").
Where("username = ? and created_at >= ? and created_at <= ?", username, startTime, endTime)
if tokenId > 0 {
query = query.Where("token_id = ?", tokenId)
}
err = query.Group("user_id, username, model_name, created_at").
Find(&quotaDatas).Error
return quotaDatas, err
}

func GetQuotaDataByUserId(userId int, startTime int64, endTime int64) (quotaData []*QuotaData, err error) {
// GetQuotaDataByUserId 根据用户 ID 查询配额数据;传入 tokenId > 0 时进一步按 API KEY 过滤。
func GetQuotaDataByUserId(userId int, startTime int64, endTime int64, tokenId int) (quotaData []*QuotaData, err error) {
var quotaDatas []*QuotaData
// 从quota_data表中查询数据
err = DB.Table("quota_data").
query := DB.Table("quota_data").
Select("user_id, username, model_name, created_at, sum(count) as count, sum(quota) as quota, sum(token_used) as token_used").
Where("user_id = ? and created_at >= ? and created_at <= ?", userId, startTime, endTime).
Group("user_id, username, model_name, created_at").
Where("user_id = ? and created_at >= ? and created_at <= ?", userId, startTime, endTime)
if tokenId > 0 {
query = query.Where("token_id = ?", tokenId)
}
err = query.Group("user_id, username, model_name, created_at").
Find(&quotaDatas).Error
return quotaDatas, err
}
Expand All @@ -170,14 +178,19 @@ func GetQuotaDataGroupByUser(startTime int64, endTime int64) (quotaData []*Quota
return quotaDatas, err
}

func GetAllQuotaDates(startTime int64, endTime int64, username string) (quotaData []*QuotaData, err error) {
// GetAllQuotaDates 查询全部配额数据;传入 username 时按用户过滤,传入 tokenId > 0 时按 API KEY 过滤。
func GetAllQuotaDates(startTime int64, endTime int64, username string, tokenId int) (quotaData []*QuotaData, err error) {
if username != "" {
return GetQuotaDataByUsername(username, startTime, endTime)
return GetQuotaDataByUsername(username, startTime, endTime, tokenId)
}
var quotaDatas []*QuotaData
// 从quota_data表中查询数据
// only select model_name, sum(count) as count, sum(quota) as quota, model_name, created_at from quota_data group by model_name, created_at;
//err = DB.Table("quota_data").Where("created_at >= ? and created_at <= ?", startTime, endTime).Find(&quotaDatas).Error
err = DB.Table("quota_data").Select("model_name, sum(count) as count, sum(quota) as quota, sum(token_used) as token_used, created_at").Where("created_at >= ? and created_at <= ?", startTime, endTime).Group("model_name, created_at").Find(&quotaDatas).Error
query := DB.Table("quota_data").
Select("model_name, sum(count) as count, sum(quota) as quota, sum(token_used) as token_used, created_at").
Where("created_at >= ? and created_at <= ?", startTime, endTime)
if tokenId > 0 {
query = query.Where("token_id = ?", tokenId)
}
err = query.Group("model_name, created_at").Find(&quotaDatas).Error
return quotaDatas, err
}
75 changes: 75 additions & 0 deletions model/usedata_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,75 @@
package model

import (
"testing"

"github.com/stretchr/testify/require"
)

func setupUsedataTestDB(t *testing.T) {
t.Helper()
truncateTables(t)
require.NoError(t, DB.Create(&Token{Id: 11, UserId: 1, Key: "sk-primary", Name: "primary"}).Error)
require.NoError(t, DB.Create(&Token{Id: 22, UserId: 2, Key: "sk-backup", Name: "backup"}).Error)
require.NoError(t, DB.Create(&QuotaData{
UserID: 1,
Username: "alice",
TokenID: 11,
ModelName: "gpt-a",
CreatedAt: 1100,
Count: 2,
Quota: 100,
TokenUsed: 40,
}).Error)
require.NoError(t, DB.Create(&QuotaData{
UserID: 2,
Username: "bob",
TokenID: 22,
ModelName: "gpt-b",
CreatedAt: 1200,
Count: 1,
Quota: 70,
TokenUsed: 30,
}).Error)
}

func TestGetAllQuotaDatesByTokenID(t *testing.T) {
setupUsedataTestDB(t)

rows, err := GetAllQuotaDates(1000, 2000, "", 11)
require.NoError(t, err)
require.Len(t, rows, 1)
require.Equal(t, "gpt-a", rows[0].ModelName)
require.Equal(t, 2, rows[0].Count)

rows, err = GetAllQuotaDates(1000, 2000, "", 0)
require.NoError(t, err)
require.Len(t, rows, 2)
}

func TestGetQuotaDataByUserIdWithTokenID(t *testing.T) {
setupUsedataTestDB(t)

rows, err := GetQuotaDataByUserId(1, 1000, 2000, 11)
require.NoError(t, err)
require.Len(t, rows, 1)
require.Equal(t, "alice", rows[0].Username)
require.Equal(t, "gpt-a", rows[0].ModelName)

rows, err = GetQuotaDataByUserId(1, 1000, 2000, 22)
require.NoError(t, err)
require.Empty(t, rows)
}

func TestGetQuotaDataByUsernameWithTokenID(t *testing.T) {
setupUsedataTestDB(t)

rows, err := GetQuotaDataByUsername("alice", 1000, 2000, 11)
require.NoError(t, err)
require.Len(t, rows, 1)
require.Equal(t, "gpt-a", rows[0].ModelName)

rows, err = GetQuotaDataByUsername("alice", 1000, 2000, 22)
require.NoError(t, err)
require.Empty(t, rows)
}
Loading