Preserve Muse relay tags without hijacking new-api request ids

Muse now forwards business context alongside each per-user token request, so new-api needs to ingest those tags for relay selection, billing audit, and log correlation without breaking its own internal request-id semantics.

This adds a dedicated Muse request-context middleware, carries Muse tags into RelayInfo, and appends them into consume/error log payloads. Muse request ids are stored as separate business metadata instead of overwriting the internal request id used for pre-consume idempotency.

Constraint: The relay layer must keep working for both /v1 and /v1beta request paths
Constraint: Client-supplied Muse request ids cannot replace the internal request id used by new-api billing/idempotency logic
Rejected: Overwriting common.RequestIdKey with X-Request-Id | mixes external business ids with internal idempotency ids
Rejected: Adding new log table columns in phase 1 | this task only needs relay/log propagation, so other JSON is the minimal compatible path
Confidence: medium
Scope-risk: moderate
Reversibility: clean
Directive: Keep Muse business identifiers in dedicated context keys and log other fields; do not reuse internal request-id slots for external correlation ids
Tested: go test ./middleware ./relay/common ./controller -run MuseRequestContext -count=1
Tested: go test ./middleware ./relay/common ./controller -count=1
Not-tested: Live relay call from Muse into a running new-api instance
This commit is contained in:
zizi 2026-04-16 19:25:29 +08:00
parent 3f24ba87ee
commit 67c173cbcb
8 changed files with 377 additions and 8 deletions

View File

@ -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"

View File

@ -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 {

View File

@ -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))
}
}

View File

@ -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"`)
}

View File

@ -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

View File

@ -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),

View File

@ -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)
}

View File

@ -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())
{