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
63 changes: 63 additions & 0 deletions controller/usedata.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,24 @@ import (
"github.com/gin-gonic/gin"
)

func parseFlowQuotaTimeRange(c *gin.Context) (int64, int64, bool) {
startTimestamp, err := strconv.ParseInt(c.Query("start_timestamp"), 10, 64)
if err != nil || startTimestamp <= 0 {
common.ApiErrorMsg(c, "invalid start_timestamp")
return 0, 0, false
}
endTimestamp, err := strconv.ParseInt(c.Query("end_timestamp"), 10, 64)
if err != nil || endTimestamp <= 0 {
common.ApiErrorMsg(c, "invalid end_timestamp")
return 0, 0, false
}
if endTimestamp < startTimestamp {
common.ApiErrorMsg(c, "invalid time range")
return 0, 0, false
}
return startTimestamp, endTimestamp, true
}

func GetAllQuotaDates(c *gin.Context) {
startTimestamp, _ := strconv.ParseInt(c.Query("start_timestamp"), 10, 64)
endTimestamp, _ := strconv.ParseInt(c.Query("end_timestamp"), 10, 64)
Expand Down Expand Up @@ -66,3 +84,48 @@ func GetUserQuotaDates(c *gin.Context) {
})
return
}

func GetAllFlowQuotaDates(c *gin.Context) {
startTimestamp, endTimestamp, ok := parseFlowQuotaTimeRange(c)
if !ok {
return
}
username := c.Query("username")
dates, err := model.GetFlowQuotaData(startTimestamp, endTimestamp, username, 0, c.GetInt("role"))
if err != nil {
common.ApiError(c, err)
return
}
c.JSON(http.StatusOK, gin.H{
"success": true,
"message": "",
"data": dates,
})
return
}
Comment on lines +88 to +105

func GetUserFlowQuotaDates(c *gin.Context) {
userId := c.GetInt("id")
startTimestamp, endTimestamp, ok := parseFlowQuotaTimeRange(c)
if !ok {
return
}
if endTimestamp-startTimestamp > 2592000 {
c.JSON(http.StatusOK, gin.H{
"success": false,
"message": "时间跨度不能超过 1 个月",
})
return
}
Comment on lines +113 to +119
dates, err := model.GetFlowQuotaData(startTimestamp, endTimestamp, "", userId, common.RoleCommonUser)
if err != nil {
common.ApiError(c, err)
return
}
c.JSON(http.StatusOK, gin.H{
"success": true,
"message": "",
"data": dates,
})
return
}
135 changes: 135 additions & 0 deletions controller/usedata_flow_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,135 @@
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 flowQuotaResponse struct {
Success bool `json:"success"`
Message string `json:"message"`
Data []model.FlowQuotaData `json:"data"`
}

func setupFlowControllerTestDB(t *testing.T) {
t.Helper()
db := setupModelListControllerTestDB(t)
require.NoError(t, db.AutoMigrate(&model.Token{}, &model.QuotaData{}))
require.NoError(t, model.DB.Create(&model.Channel{Id: 1, Name: "east"}).Error)
require.NoError(t, model.DB.Create(&model.Token{Id: 11, UserId: 1, Key: "sk-primary", Name: "primary"}).Error)
require.NoError(t, model.DB.Create(&model.Token{Id: 22, UserId: 2, Key: "sk-backup", Name: "backup"}).Error)
require.NoError(t, model.DB.Create(&model.QuotaData{
UserID: 1,
Username: "alice",
NodeName: "node-a",
TokenID: 11,
UseGroup: "default",
ChannelID: 1,
ModelName: "gpt-a",
CreatedAt: 1100,
Count: 2,
Quota: 100,
TokenUsed: 40,
}).Error)
require.NoError(t, model.DB.Create(&model.QuotaData{
UserID: 2,
Username: "bob",
NodeName: "node-b",
TokenID: 22,
UseGroup: "vip",
ChannelID: 1,
ModelName: "gpt-b",
CreatedAt: 1200,
Count: 1,
Quota: 70,
TokenUsed: 30,
}).Error)
}

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

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

recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
ctx.Set("role", common.RoleAdminUser)
ctx.Request = httptest.NewRequest(http.MethodGet, "/api/data/flow?start_timestamp=1000&end_timestamp=2000&username=bob", nil)

GetAllFlowQuotaDates(ctx)

payload := decodeFlowQuotaResponse(t, recorder)
require.Len(t, payload.Data, 1)
require.Equal(t, "bob", payload.Data[0].Username)
require.Equal(t, "vip", payload.Data[0].UseGroup)
require.Equal(t, "east", payload.Data[0].ChannelName)
require.Empty(t, payload.Data[0].TokenName)
require.Empty(t, payload.Data[0].NodeName)
}

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

recorder := httptest.NewRecorder()
ctx, _ := gin.CreateTestContext(recorder)
ctx.Set("role", common.RoleRootUser)
ctx.Request = httptest.NewRequest(http.MethodGet, "/api/data/flow?start_timestamp=1000&end_timestamp=2000&username=alice", nil)

GetAllFlowQuotaDates(ctx)

payload := decodeFlowQuotaResponse(t, recorder)
require.Len(t, payload.Data, 1)
require.Equal(t, "alice", payload.Data[0].Username)
require.Equal(t, "node-a", payload.Data[0].NodeName)
require.Equal(t, "primary", payload.Data[0].TokenName)
require.Equal(t, "default", payload.Data[0].UseGroup)
require.Equal(t, "east", payload.Data[0].ChannelName)
}

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

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

GetUserFlowQuotaDates(ctx)

payload := decodeFlowQuotaResponse(t, recorder)
require.Len(t, payload.Data, 1)
require.Empty(t, payload.Data[0].Username)
require.Equal(t, "primary", payload.Data[0].TokenName)
require.Equal(t, "default", payload.Data[0].UseGroup)
require.Empty(t, payload.Data[0].ChannelName)
}

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

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

GetUserFlowQuotaDates(ctx)

require.Equal(t, http.StatusOK, recorder.Code)
var payload flowQuotaResponse
require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &payload))
require.False(t, payload.Success)
require.Equal(t, "invalid start_timestamp", payload.Message)
}
34 changes: 31 additions & 3 deletions model/log.go
Original file line number Diff line number Diff line change
Expand Up @@ -298,6 +298,7 @@ func RecordConsumeLog(c *gin.Context, userId int, params RecordConsumeLogParams)
username := c.GetString("username")
requestId := c.GetString(common.RequestIdKey)
upstreamRequestId := c.GetString(common.UpstreamRequestIdKey)
createdAt := common.GetTimestamp()
otherStr := common.MapToJsonStr(params.Other)
// 判断是否需要记录 IP
needRecordIp := false
Expand All @@ -309,7 +310,7 @@ func RecordConsumeLog(c *gin.Context, userId int, params RecordConsumeLogParams)
log := &Log{
UserId: userId,
Username: username,
CreatedAt: common.GetTimestamp(),
CreatedAt: createdAt,
Type: LogTypeConsume,
Content: params.Content,
PromptTokens: params.PromptTokens,
Expand Down Expand Up @@ -338,7 +339,18 @@ func RecordConsumeLog(c *gin.Context, userId int, params RecordConsumeLogParams)
}
if common.DataExportEnabled {
gopool.Go(func() {
LogQuotaData(userId, username, params.ModelName, params.Quota, common.GetTimestamp(), params.PromptTokens+params.CompletionTokens)
LogQuotaData(QuotaDataLogParams{
UserID: userId,
Username: username,
ModelName: params.ModelName,
Quota: params.Quota,
CreatedAt: createdAt,
TokenUsed: params.PromptTokens + params.CompletionTokens,
UseGroup: params.Group,
TokenID: params.TokenId,
ChannelID: params.ChannelId,
NodeName: common.NodeName,
})
})
}
}
Expand Down Expand Up @@ -366,10 +378,11 @@ func RecordTaskBillingLog(params RecordTaskBillingLogParams) {
tokenName = token.Name
}
}
createdAt := common.GetTimestamp()
log := &Log{
UserId: params.UserId,
Username: username,
CreatedAt: common.GetTimestamp(),
CreatedAt: createdAt,
Type: params.LogType,
Content: params.Content,
TokenName: tokenName,
Expand All @@ -384,6 +397,21 @@ func RecordTaskBillingLog(params RecordTaskBillingLogParams) {
if err != nil {
common.SysLog("failed to record task billing log: " + err.Error())
}
if params.LogType == LogTypeConsume && common.DataExportEnabled {
gopool.Go(func() {
LogQuotaData(QuotaDataLogParams{
UserID: params.UserId,
Username: username,
ModelName: params.ModelName,
Quota: params.Quota,
CreatedAt: createdAt,
UseGroup: params.Group,
TokenID: params.TokenId,
ChannelID: params.ChannelId,
NodeName: common.NodeName,
})
})
}
}

func GetAllLogs(logType int, startTimestamp int64, endTimestamp int64, modelName string, username string, tokenName string, startIdx int, num int, channel int, group string, requestId string, upstreamRequestId string) (logs []*Log, total int64, err error) {
Expand Down
2 changes: 2 additions & 0 deletions model/task_cas_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@ func TestMain(m *testing.M) {
&Token{},
&Log{},
&Channel{},
&QuotaData{},
&Ability{},
&TopUp{},
&SubscriptionPlan{},
Expand All @@ -62,6 +63,7 @@ func truncateTables(t *testing.T) {
DB.Exec("DELETE FROM tokens")
DB.Exec("DELETE FROM logs")
DB.Exec("DELETE FROM channels")
DB.Exec("DELETE FROM quota_data")
DB.Exec("DELETE FROM abilities")
DB.Exec("DELETE FROM top_ups")
DB.Exec("DELETE FROM subscription_orders")
Expand Down
Loading