You can not select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
 
 
 

248 lines
7.0 KiB

  1. package model
  2. import (
  3. "net/http/httptest"
  4. "testing"
  5. "github.com/QuantumNous/new-api/common"
  6. relaycommon "github.com/QuantumNous/new-api/relay/common"
  7. "github.com/gin-gonic/gin"
  8. "github.com/glebarez/sqlite"
  9. "github.com/stretchr/testify/require"
  10. "gorm.io/gorm"
  11. )
  12. func setupLogIdentityDB(t *testing.T) *gorm.DB {
  13. t.Helper()
  14. db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
  15. require.NoError(t, err)
  16. origDB := DB
  17. origLogDB := LOG_DB
  18. origSQLite := common.UsingSQLite
  19. origMySQL := common.UsingMySQL
  20. origPostgres := common.UsingPostgreSQL
  21. DB = db
  22. LOG_DB = db
  23. common.UsingSQLite = true
  24. common.UsingMySQL = false
  25. common.UsingPostgreSQL = false
  26. initCol()
  27. require.NoError(t, db.AutoMigrate(&Log{}))
  28. t.Cleanup(func() {
  29. DB = origDB
  30. LOG_DB = origLogDB
  31. common.UsingSQLite = origSQLite
  32. common.UsingMySQL = origMySQL
  33. common.UsingPostgreSQL = origPostgres
  34. initCol()
  35. })
  36. return db
  37. }
  38. func TestLogPersistsChatIDAndUpstreamID(t *testing.T) {
  39. db := setupLogIdentityDB(t)
  40. entry := &Log{
  41. UserId: 1,
  42. Username: "alice",
  43. CreatedAt: 1714464000,
  44. Type: LogTypeConsume,
  45. Content: "usage",
  46. ModelName: "gpt-4o-mini",
  47. TokenName: "demo",
  48. ChatId: "chat_123",
  49. UpstreamId: "req_upstream_123",
  50. RequestId: "req_internal_123",
  51. }
  52. require.NoError(t, db.Create(entry).Error)
  53. var saved Log
  54. require.NoError(t, db.First(&saved).Error)
  55. require.Equal(t, "chat_123", saved.ChatId)
  56. require.Equal(t, "req_upstream_123", saved.UpstreamId)
  57. require.Equal(t, "req_internal_123", saved.RequestId)
  58. }
  59. func TestGetAllLogsFiltersByChatIDAndUpstreamID(t *testing.T) {
  60. db := setupLogIdentityDB(t)
  61. require.NoError(t, db.Create(&Log{
  62. UserId: 1,
  63. Username: "alice",
  64. CreatedAt: 1714464001,
  65. Type: LogTypeConsume,
  66. Content: "usage-a",
  67. ModelName: "gpt-4o-mini",
  68. TokenName: "demo",
  69. ChatId: "chat_a",
  70. UpstreamId: "up_a",
  71. RequestId: "req_a",
  72. }).Error)
  73. require.NoError(t, db.Create(&Log{
  74. UserId: 2,
  75. Username: "bob",
  76. CreatedAt: 1714464002,
  77. Type: LogTypeConsume,
  78. Content: "usage-b",
  79. ModelName: "gpt-4o-mini",
  80. TokenName: "demo",
  81. ChatId: "chat_b",
  82. UpstreamId: "up_b",
  83. RequestId: "req_b",
  84. }).Error)
  85. logs, total, err := GetAllLogs(LogTypeConsume, 0, 0, "", "", "", 0, 20, 0, "", "", "chat_a", "")
  86. require.NoError(t, err)
  87. require.Equal(t, int64(1), total)
  88. require.Len(t, logs, 1)
  89. require.Equal(t, "chat_a", logs[0].ChatId)
  90. logs, total, err = GetAllLogs(LogTypeConsume, 0, 0, "", "", "", 0, 20, 0, "", "", "", "up_b")
  91. require.NoError(t, err)
  92. require.Equal(t, int64(1), total)
  93. require.Len(t, logs, 1)
  94. require.Equal(t, "up_b", logs[0].UpstreamId)
  95. }
  96. func TestGetUserLogsFiltersByChatIDAndUpstreamID(t *testing.T) {
  97. db := setupLogIdentityDB(t)
  98. require.NoError(t, db.Create(&Log{
  99. UserId: 9,
  100. Username: "owner",
  101. CreatedAt: 1714464010,
  102. Type: LogTypeConsume,
  103. Content: "usage-owner-a",
  104. ModelName: "gpt-4o-mini",
  105. TokenName: "demo",
  106. ChatId: "chat-owner-a",
  107. UpstreamId: "up-owner-a",
  108. RequestId: "req-owner-a",
  109. }).Error)
  110. require.NoError(t, db.Create(&Log{
  111. UserId: 9,
  112. Username: "owner",
  113. CreatedAt: 1714464011,
  114. Type: LogTypeConsume,
  115. Content: "usage-owner-b",
  116. ModelName: "gpt-4o-mini",
  117. TokenName: "demo",
  118. ChatId: "chat-owner-b",
  119. UpstreamId: "up-owner-b",
  120. RequestId: "req-owner-b",
  121. }).Error)
  122. logs, total, err := GetUserLogs(9, LogTypeConsume, 0, 0, "", "", 0, 20, "", "", "chat-owner-b", "")
  123. require.NoError(t, err)
  124. require.Equal(t, int64(1), total)
  125. require.Len(t, logs, 1)
  126. require.Equal(t, "chat-owner-b", logs[0].ChatId)
  127. logs, total, err = GetUserLogs(9, LogTypeConsume, 0, 0, "", "", 0, 20, "", "", "", "up-owner-a")
  128. require.NoError(t, err)
  129. require.Equal(t, int64(1), total)
  130. require.Len(t, logs, 1)
  131. require.Equal(t, "up-owner-a", logs[0].UpstreamId)
  132. }
  133. func TestGetUserLogsIncludesTaskSettlementOutsidePage(t *testing.T) {
  134. db := setupLogIdentityDB(t)
  135. preconsume := &Log{
  136. UserId: 1, CreatedAt: 100, Type: LogTypeConsume, ModelName: "seedance",
  137. Other: `{"billing_mode":"matrix","billing_phase":"preconsume","task_id":"task_cross_page"}`,
  138. }
  139. settlement := &Log{
  140. UserId: 1, CreatedAt: 200, Type: LogTypeConsume, ModelName: "seedance",
  141. Other: `{"billing_mode":"matrix","billing_phase":"settlement","task_id":"task_cross_page"}`,
  142. }
  143. require.NoError(t, db.Create(preconsume).Error)
  144. require.NoError(t, db.Create(settlement).Error)
  145. logs, total, err := GetUserLogs(1, LogTypeConsume, 0, 0, "", "", 1, 1, "", "", "", "")
  146. require.NoError(t, err)
  147. require.Equal(t, int64(2), total)
  148. require.Len(t, logs, 2)
  149. require.Equal(t, "preconsume", logBillingPhase(t, logs[0]))
  150. require.Equal(t, "settlement", logBillingPhase(t, logs[1]))
  151. }
  152. func TestGetAllLogsIncludesTaskSettlementOutsidePage(t *testing.T) {
  153. db := setupLogIdentityDB(t)
  154. preconsume := &Log{
  155. UserId: 1, CreatedAt: 100, Type: LogTypeConsume, ModelName: "seedance",
  156. Other: `{"billing_mode":"matrix","billing_phase":"preconsume","task_id":"task_cross_page_admin"}`,
  157. }
  158. settlement := &Log{
  159. UserId: 1, CreatedAt: 200, Type: LogTypeConsume, ModelName: "seedance",
  160. Other: `{"billing_mode":"matrix","billing_phase":"settlement","task_id":"task_cross_page_admin"}`,
  161. }
  162. require.NoError(t, db.Create(preconsume).Error)
  163. require.NoError(t, db.Create(settlement).Error)
  164. logs, total, err := GetAllLogs(LogTypeConsume, 0, 0, "", "", "", 1, 1, 0, "", "", "", "")
  165. require.NoError(t, err)
  166. require.Equal(t, int64(2), total)
  167. require.Len(t, logs, 2)
  168. require.Equal(t, "preconsume", logBillingPhase(t, logs[0]))
  169. require.Equal(t, "settlement", logBillingPhase(t, logs[1]))
  170. }
  171. func logBillingPhase(t *testing.T, log *Log) string {
  172. t.Helper()
  173. other := map[string]any{}
  174. require.NoError(t, common.UnmarshalJsonStr(log.Other, &other))
  175. phase, _ := other["billing_phase"].(string)
  176. return phase
  177. }
  178. func TestRecordConsumeLogPersistsContextChatIDAndUpstreamID(t *testing.T) {
  179. db := setupLogIdentityDB(t)
  180. gin.SetMode(gin.TestMode)
  181. w := httptest.NewRecorder()
  182. c, _ := gin.CreateTestContext(w)
  183. c.Set("username", "alice")
  184. c.Set(common.RequestIdKey, "req_internal_ctx")
  185. relaycommon.SetRelayChatID(c, "chat_ctx")
  186. relaycommon.SetRelayUpstreamID(c, "up_ctx")
  187. RecordConsumeLog(c, 1, RecordConsumeLogParams{
  188. ChannelId: 1,
  189. ModelName: "gpt-4o-mini",
  190. TokenName: "demo",
  191. Content: "usage",
  192. Quota: 10,
  193. })
  194. var saved Log
  195. require.NoError(t, db.First(&saved).Error)
  196. require.Equal(t, "chat_ctx", saved.ChatId)
  197. require.Equal(t, "up_ctx", saved.UpstreamId)
  198. }
  199. func TestRecordErrorLogPersistsContextChatIDAndUpstreamID(t *testing.T) {
  200. db := setupLogIdentityDB(t)
  201. gin.SetMode(gin.TestMode)
  202. w := httptest.NewRecorder()
  203. c, _ := gin.CreateTestContext(w)
  204. c.Set("username", "alice")
  205. c.Set(common.RequestIdKey, "req_internal_err")
  206. relaycommon.SetRelayChatID(c, "chat_err")
  207. relaycommon.SetRelayUpstreamID(c, "up_err")
  208. RecordErrorLog(c, 1, 2, "gpt-4o-mini", "demo", "upstream failed", 3, 4, false, "default", map[string]interface{}{})
  209. var saved Log
  210. require.NoError(t, db.First(&saved).Error)
  211. require.Equal(t, "chat_err", saved.ChatId)
  212. require.Equal(t, "up_err", saved.UpstreamId)
  213. }