diff --git a/constant/context_key.go b/constant/context_key.go index 2ba2fe27..cb4a1e3c 100644 --- a/constant/context_key.go +++ b/constant/context_key.go @@ -56,6 +56,12 @@ const ( ContextKeySystemPromptOverride ContextKey = "system_prompt_override" + ContextKeyMuseUserID ContextKey = "muse_user_id" + ContextKeyMuseWorkID ContextKey = "muse_work_id" + ContextKeyMuseRequestID ContextKey = "muse_request_id" + ContextKeyMuseScene ContextKey = "muse_scene" + ContextKeyMuseTraceID ContextKey = "muse_trace_id" + // ContextKeyFileSourcesToCleanup stores file sources that need cleanup when request ends ContextKeyFileSourcesToCleanup ContextKey = "file_sources_to_cleanup" diff --git a/controller/log.go b/controller/log.go index cf3825f1..a5088733 100644 --- a/controller/log.go +++ b/controller/log.go @@ -20,6 +20,7 @@ func GetAllLogs(c *gin.Context) { modelName := c.Query("model_name") channel, _ := strconv.Atoi(c.Query("channel")) group := c.Query("group") + // request_id filters the internal relay request id; muse_request_id is carried in log.other for drill-down. requestId := c.Query("request_id") logs, total, err := model.GetAllLogs(logType, startTimestamp, endTimestamp, modelName, username, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), channel, group, requestId) if err != nil { @@ -41,6 +42,7 @@ func GetUserLogs(c *gin.Context) { tokenName := c.Query("token_name") modelName := c.Query("model_name") group := c.Query("group") + // request_id filters the internal relay request id; muse_request_id is carried in log.other for drill-down. requestId := c.Query("request_id") logs, total, err := model.GetUserLogs(userId, logType, startTimestamp, endTimestamp, modelName, tokenName, pageInfo.GetStartIdx(), pageInfo.GetPageSize(), group, requestId) if err != nil { diff --git a/middleware/muse_request_context.go b/middleware/muse_request_context.go new file mode 100644 index 00000000..3d1cdd17 --- /dev/null +++ b/middleware/muse_request_context.go @@ -0,0 +1,113 @@ +package middleware + +import ( + "fmt" + "strings" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/constant" + "github.com/gin-gonic/gin" +) + +const ( + museHeaderUserID = "X-Muse-User-Id" + museHeaderWorkID = "X-Muse-Work-Id" + museHeaderRequest = "X-Request-Id" + museHeaderTraceID = "X-Trace-Id" + museHeaderScene = "X-Muse-Scene" + jsonContentType = "application/json" +) + +type museMetadataEnvelope struct { + Metadata map[string]any `json:"metadata"` +} + +func MuseRequestContext() gin.HandlerFunc { + return func(c *gin.Context) { + metadata := extractMuseMetadata(c) + + museUserID := firstNonBlank( + c.GetHeader(museHeaderUserID), + stringifyMetadata(metadata["muse_user_id"]), + ) + museWorkID := firstNonBlank( + c.GetHeader(museHeaderWorkID), + stringifyMetadata(metadata["work_id"]), + stringifyMetadata(metadata["muse_work_id"]), + ) + museScene := firstNonBlank( + c.GetHeader(museHeaderScene), + stringifyMetadata(metadata["scene"]), + ) + requestID := firstNonBlank( + c.GetHeader(museHeaderRequest), + stringifyMetadata(metadata["request_id"]), + ) + traceID := firstNonBlank( + c.GetHeader(museHeaderTraceID), + stringifyMetadata(metadata["trace_id"]), + ) + + if museUserID != "" { + common.SetContextKey(c, constant.ContextKeyMuseUserID, museUserID) + } + if museWorkID != "" { + common.SetContextKey(c, constant.ContextKeyMuseWorkID, museWorkID) + } + if requestID != "" { + common.SetContextKey(c, constant.ContextKeyMuseRequestID, requestID) + } + if museScene != "" { + common.SetContextKey(c, constant.ContextKeyMuseScene, museScene) + } + if traceID != "" { + common.SetContextKey(c, constant.ContextKeyMuseTraceID, traceID) + } + c.Next() + } +} + +func extractMuseMetadata(c *gin.Context) map[string]any { + if c == nil || c.Request == nil { + return nil + } + if !strings.Contains(strings.ToLower(c.GetHeader("Content-Type")), jsonContentType) { + return nil + } + if !needsMuseMetadataRead(c) { + return nil + } + var envelope museMetadataEnvelope + if err := common.UnmarshalBodyReusable(c, &envelope); err != nil { + return nil + } + return envelope.Metadata +} + +func needsMuseMetadataRead(c *gin.Context) bool { + return c.GetHeader(museHeaderUserID) == "" || + c.GetHeader(museHeaderWorkID) == "" || + c.GetHeader(museHeaderRequest) == "" || + c.GetHeader(museHeaderTraceID) == "" || + c.GetHeader(museHeaderScene) == "" +} + +func firstNonBlank(values ...string) string { + for _, value := range values { + if strings.TrimSpace(value) != "" { + return strings.TrimSpace(value) + } + } + return "" +} + +func stringifyMetadata(value any) string { + switch v := value.(type) { + case nil: + return "" + case string: + return strings.TrimSpace(v) + default: + return strings.TrimSpace(fmt.Sprintf("%v", v)) + } +} diff --git a/middleware/muse_request_context_test.go b/middleware/muse_request_context_test.go new file mode 100644 index 00000000..2ba4025b --- /dev/null +++ b/middleware/muse_request_context_test.go @@ -0,0 +1,176 @@ +package middleware + +import ( + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/constant" + "github.com/QuantumNous/new-api/model" + "github.com/gin-gonic/gin" + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +type museRequestContextResponse struct { + MuseUserID string `json:"muse_user_id"` + MuseWorkID string `json:"muse_work_id"` + MuseScene string `json:"muse_scene"` + MuseRequestID string `json:"muse_request_id"` + TraceID string `json:"trace_id"` +} + +func setupMuseRequestContextTestDB(t *testing.T) *gorm.DB { + t.Helper() + + common.UsingSQLite = true + common.UsingMySQL = false + common.UsingPostgreSQL = false + common.RedisEnabled = false + + db, err := gorm.Open(sqlite.Open("file:muse_request_context_test?mode=memory&cache=shared"), &gorm.Config{}) + require.NoError(t, err) + model.DB = db + model.LOG_DB = db + require.NoError(t, db.AutoMigrate(&model.User{}, &model.Log{})) + require.NoError(t, db.Create(&model.User{ + Id: 1, + Username: "alice", + Password: "hashed", + Status: common.UserStatusEnabled, + }).Error) + + t.Cleanup(func() { + sqlDB, err := db.DB() + if err == nil { + _ = sqlDB.Close() + } + }) + return db +} + +func TestMuseRequestContextExtractsHeaders(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest( + http.MethodPost, + "/v1/chat/completions", + strings.NewReader(`{"metadata":{"work_id":"42","scene":"suggestion_generation","muse_user_id":"u-1"}}`), + ) + req.Header.Set("Content-Type", "application/json") + req.Header.Set(museHeaderUserID, "u-1") + req.Header.Set(museHeaderWorkID, "42") + req.Header.Set(museHeaderRequest, "req-1") + req.Header.Set(museHeaderTraceID, "trace-1") + + recorder := httptest.NewRecorder() + router := gin.New() + router.Use(MuseRequestContext()) + router.POST("/v1/chat/completions", func(c *gin.Context) { + c.JSON(http.StatusOK, museRequestContextResponse{ + MuseUserID: common.GetContextKeyString(c, constant.ContextKeyMuseUserID), + MuseWorkID: common.GetContextKeyString(c, constant.ContextKeyMuseWorkID), + MuseScene: common.GetContextKeyString(c, constant.ContextKeyMuseScene), + MuseRequestID: common.GetContextKeyString(c, constant.ContextKeyMuseRequestID), + TraceID: common.GetContextKeyString(c, constant.ContextKeyMuseTraceID), + }) + }) + + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusOK, recorder.Code) + + var response museRequestContextResponse + require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &response)) + require.Equal(t, "u-1", response.MuseUserID) + require.Equal(t, "42", response.MuseWorkID) + require.Equal(t, "suggestion_generation", response.MuseScene) + require.Equal(t, "trace-1", response.TraceID) + require.Equal(t, "req-1", response.MuseRequestID) +} + +func TestMuseRequestContextFallsBackToMetadata(t *testing.T) { + gin.SetMode(gin.TestMode) + + req := httptest.NewRequest( + http.MethodPost, + "/v1/chat/completions", + strings.NewReader(`{"metadata":{"muse_user_id":"u-2","work_id":42,"scene":"suggestion_generation","request_id":"req-meta","trace_id":"trace-meta"}}`), + ) + req.Header.Set("Content-Type", "application/json") + + recorder := httptest.NewRecorder() + router := gin.New() + router.Use(MuseRequestContext()) + router.POST("/v1/chat/completions", func(c *gin.Context) { + c.JSON(http.StatusOK, museRequestContextResponse{ + MuseUserID: common.GetContextKeyString(c, constant.ContextKeyMuseUserID), + MuseWorkID: common.GetContextKeyString(c, constant.ContextKeyMuseWorkID), + MuseScene: common.GetContextKeyString(c, constant.ContextKeyMuseScene), + MuseRequestID: common.GetContextKeyString(c, constant.ContextKeyMuseRequestID), + TraceID: common.GetContextKeyString(c, constant.ContextKeyMuseTraceID), + }) + }) + + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusOK, recorder.Code) + + var response museRequestContextResponse + require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &response)) + require.Equal(t, "u-2", response.MuseUserID) + require.Equal(t, "42", response.MuseWorkID) + require.Equal(t, "suggestion_generation", response.MuseScene) + require.Equal(t, "trace-meta", response.TraceID) + require.Equal(t, "req-meta", response.MuseRequestID) +} + +func TestMuseRequestContextRecordConsumeLogIncludesMuseFields(t *testing.T) { + gin.SetMode(gin.TestMode) + setupMuseRequestContextTestDB(t) + + req := httptest.NewRequest( + http.MethodPost, + "/v1/chat/completions", + strings.NewReader(`{"metadata":{"scene":"suggestion_generation","request_id":"req-log","trace_id":"trace-log"}}`), + ) + req.Header.Set("Content-Type", "application/json") + req.Header.Set(museHeaderUserID, "u-log") + req.Header.Set(museHeaderWorkID, "77") + recorder := httptest.NewRecorder() + + router := gin.New() + router.Use(MuseRequestContext()) + router.POST("/v1/chat/completions", func(c *gin.Context) { + model.RecordConsumeLog(c, 1, model.RecordConsumeLogParams{ + ModelName: "gpt-4o", + Content: "ok", + Group: "default", + Other: map[string]interface{}{"existing": true}, + }) + var logEntry model.Log + require.NoError(t, model.LOG_DB.Last(&logEntry).Error) + c.JSON(http.StatusOK, gin.H{ + "request_id": logEntry.RequestId, + "other": logEntry.Other, + }) + }) + + router.ServeHTTP(recorder, req) + + require.Equal(t, http.StatusOK, recorder.Code) + var response struct { + RequestID string `json:"request_id"` + Other string `json:"other"` + } + require.NoError(t, common.Unmarshal(recorder.Body.Bytes(), &response)) + require.Empty(t, response.RequestID) + require.Contains(t, response.Other, `"muse_user_id":"u-log"`) + require.Contains(t, response.Other, `"muse_work_id":"77"`) + require.Contains(t, response.Other, `"muse_request_id":"req-log"`) + require.Contains(t, response.Other, `"muse_scene":"suggestion_generation"`) + require.Contains(t, response.Other, `"muse_trace_id":"trace-log"`) +} diff --git a/model/log.go b/model/log.go index 68bc6504..15d7eb29 100644 --- a/model/log.go +++ b/model/log.go @@ -7,6 +7,7 @@ import ( "time" "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/constant" "github.com/QuantumNous/new-api/logger" "github.com/QuantumNous/new-api/types" @@ -95,6 +96,7 @@ func RecordErrorLog(c *gin.Context, userId int, channelId int, modelName string, logger.LogInfo(c, fmt.Sprintf("record error log: userId=%d, channelId=%d, modelName=%s, tokenName=%s, content=%s", userId, channelId, modelName, tokenName, content)) username := c.GetString("username") requestId := c.GetString(common.RequestIdKey) + other = appendMuseContextToOther(c, other) otherStr := common.MapToJsonStr(other) // 判断是否需要记录 IP needRecordIp := false @@ -156,6 +158,7 @@ func RecordConsumeLog(c *gin.Context, userId int, params RecordConsumeLogParams) logger.LogInfo(c, fmt.Sprintf("record consume log: userId=%d, params=%s", userId, common.GetJsonString(params))) username := c.GetString("username") requestId := c.GetString(common.RequestIdKey) + params.Other = appendMuseContextToOther(c, params.Other) otherStr := common.MapToJsonStr(params.Other) // 判断是否需要记录 IP needRecordIp := false @@ -200,6 +203,31 @@ func RecordConsumeLog(c *gin.Context, userId int, params RecordConsumeLogParams) } } +func appendMuseContextToOther(c *gin.Context, other map[string]interface{}) map[string]interface{} { + if c == nil { + return other + } + if other == nil { + other = make(map[string]interface{}) + } + if museUserID := common.GetContextKeyString(c, constant.ContextKeyMuseUserID); museUserID != "" { + other["muse_user_id"] = museUserID + } + if museWorkID := common.GetContextKeyString(c, constant.ContextKeyMuseWorkID); museWorkID != "" { + other["muse_work_id"] = museWorkID + } + if museRequestID := common.GetContextKeyString(c, constant.ContextKeyMuseRequestID); museRequestID != "" { + other["muse_request_id"] = museRequestID + } + if museScene := common.GetContextKeyString(c, constant.ContextKeyMuseScene); museScene != "" { + other["muse_scene"] = museScene + } + if museTraceID := common.GetContextKeyString(c, constant.ContextKeyMuseTraceID); museTraceID != "" { + other["muse_trace_id"] = museTraceID + } + return other +} + type RecordTaskBillingLogParams struct { UserId int LogType int diff --git a/relay/common/relay_info.go b/relay/common/relay_info.go index e4421fc1..c5739388 100644 --- a/relay/common/relay_info.go +++ b/relay/common/relay_info.go @@ -139,7 +139,12 @@ type RelayInfo struct { SubscriptionPlanId int SubscriptionPlanTitle string // RequestId is used for idempotent pre-consume/refund - RequestId string + RequestId string + MuseUserID string + MuseWorkID string + MuseRequestID string + MuseScene string + MuseTraceID string // SubscriptionAmountTotal / SubscriptionAmountUsedAfterPreConsume are used to compute remaining in logs. SubscriptionAmountTotal int64 SubscriptionAmountUsedAfterPreConsume int64 @@ -255,7 +260,6 @@ func (info *RelayInfo) ToString() string { latencyMs := info.FirstResponseTime.Sub(info.StartTime).Milliseconds() fmt.Fprintf(b, "Timing{ Start: %s, FirstResponse: %s, LatencyMs: %d }, ", info.StartTime.Format(time.RFC3339Nano), info.FirstResponseTime.Format(time.RFC3339Nano), latencyMs) - // Audio / realtime if info.InputAudioFormat != "" || info.OutputAudioFormat != "" || len(info.RealtimeTools) > 0 || info.AudioUsage { fmt.Fprintf(b, "Realtime{ AudioUsage: %t, InFmt: %q, OutFmt: %q, Tools: %d }, ", @@ -447,12 +451,17 @@ func genBaseRelayInfo(c *gin.Context, request dto.Request) *RelayInfo { info := &RelayInfo{ Request: request, - RequestId: reqId, - UserId: common.GetContextKeyInt(c, constant.ContextKeyUserId), - UsingGroup: common.GetContextKeyString(c, constant.ContextKeyUsingGroup), - UserGroup: common.GetContextKeyString(c, constant.ContextKeyUserGroup), - UserQuota: common.GetContextKeyInt(c, constant.ContextKeyUserQuota), - UserEmail: common.GetContextKeyString(c, constant.ContextKeyUserEmail), + RequestId: reqId, + MuseUserID: common.GetContextKeyString(c, constant.ContextKeyMuseUserID), + MuseWorkID: common.GetContextKeyString(c, constant.ContextKeyMuseWorkID), + MuseRequestID: common.GetContextKeyString(c, constant.ContextKeyMuseRequestID), + MuseScene: common.GetContextKeyString(c, constant.ContextKeyMuseScene), + MuseTraceID: common.GetContextKeyString(c, constant.ContextKeyMuseTraceID), + UserId: common.GetContextKeyInt(c, constant.ContextKeyUserId), + UsingGroup: common.GetContextKeyString(c, constant.ContextKeyUsingGroup), + UserGroup: common.GetContextKeyString(c, constant.ContextKeyUserGroup), + UserQuota: common.GetContextKeyInt(c, constant.ContextKeyUserQuota), + UserEmail: common.GetContextKeyString(c, constant.ContextKeyUserEmail), OriginModelName: common.GetContextKeyString(c, constant.ContextKeyOriginalModel), diff --git a/relay/common/relay_info_test.go b/relay/common/relay_info_test.go index e53ec804..2c77319a 100644 --- a/relay/common/relay_info_test.go +++ b/relay/common/relay_info_test.go @@ -1,9 +1,15 @@ package common import ( + "net/http" + "net/http/httptest" "testing" + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/constant" + "github.com/QuantumNous/new-api/dto" "github.com/QuantumNous/new-api/types" + "github.com/gin-gonic/gin" "github.com/stretchr/testify/require" ) @@ -38,3 +44,30 @@ func TestRelayInfoGetFinalRequestRelayFormatNilReceiver(t *testing.T) { var info *RelayInfo require.Equal(t, types.RelayFormat(""), info.GetFinalRequestRelayFormat()) } + +func TestMuseRequestContextGenRelayInfoCopiesMuseFields(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + ctx, _ := gin.CreateTestContext(recorder) + ctx.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil) + + ctx.Set(common.RequestIdKey, "req-1") + common.SetContextKey(ctx, constant.ContextKeyUserId, 11) + common.SetContextKey(ctx, constant.ContextKeyUserGroup, "default") + common.SetContextKey(ctx, constant.ContextKeyUsingGroup, "default") + common.SetContextKey(ctx, constant.ContextKeyMuseUserID, "muse-user-1") + common.SetContextKey(ctx, constant.ContextKeyMuseWorkID, "42") + common.SetContextKey(ctx, constant.ContextKeyMuseRequestID, "muse-req-1") + common.SetContextKey(ctx, constant.ContextKeyMuseScene, "suggestion_generation") + common.SetContextKey(ctx, constant.ContextKeyMuseTraceID, "trace-1") + + info := GenRelayInfoOpenAI(ctx, &dto.GeneralOpenAIRequest{}) + + require.NotNil(t, info) + require.Equal(t, "req-1", info.RequestId) + require.Equal(t, "muse-user-1", info.MuseUserID) + require.Equal(t, "42", info.MuseWorkID) + require.Equal(t, "muse-req-1", info.MuseRequestID) + require.Equal(t, "suggestion_generation", info.MuseScene) + require.Equal(t, "trace-1", info.MuseTraceID) +} diff --git a/router/relay-router.go b/router/relay-router.go index 17a13cad..8dc6fab1 100644 --- a/router/relay-router.go +++ b/router/relay-router.go @@ -70,6 +70,7 @@ func SetRelayRouter(router *gin.Engine) { relayV1Router.Use(middleware.RouteTag("relay")) relayV1Router.Use(middleware.SystemPerformanceCheck()) relayV1Router.Use(middleware.TokenAuth()) + relayV1Router.Use(middleware.MuseRequestContext()) relayV1Router.Use(middleware.ModelRequestRateLimit()) { // WebSocket 路由(统一到 Relay) @@ -190,6 +191,7 @@ func SetRelayRouter(router *gin.Engine) { relayGeminiRouter.Use(middleware.RouteTag("relay")) relayGeminiRouter.Use(middleware.SystemPerformanceCheck()) relayGeminiRouter.Use(middleware.TokenAuth()) + relayGeminiRouter.Use(middleware.MuseRequestContext()) relayGeminiRouter.Use(middleware.ModelRequestRateLimit()) relayGeminiRouter.Use(middleware.Distribute()) {