feat(usage): persist client session identifiers

This commit is contained in:
Edison42
2026-07-24 01:22:34 +08:00
parent cb24522dd5
commit 1c0cb24c7e
33 changed files with 530 additions and 21 deletions
@@ -42,6 +42,9 @@ func (h *BatchImageHandler) Submit(c *gin.Context) {
if !h.checkSecurityAuditBeforeSubmit(c, &req) {
return
}
if sessionID := service.ExtractClientSessionID(c); sessionID != "" {
req.SessionID = &sessionID
}
got, err := h.service.Submit(c.Request.Context(), owner, req, c.GetHeader("Idempotency-Key"))
if err != nil {
batchImageError(c, err)
+1
View File
@@ -669,6 +669,7 @@ func usageLogFromServiceUser(l *service.UsageLog) UsageLog {
MediaType: l.MediaType,
UserAgent: l.UserAgent,
IPAddress: l.IPAddress,
SessionID: l.SessionID,
CacheTTLOverridden: l.CacheTTLOverridden,
BillingMode: l.BillingMode,
CreatedAt: l.CreatedAt,
+3
View File
@@ -521,6 +521,9 @@ type UsageLog struct {
UserAgent *string `json:"user_agent"`
// IPAddress is visible to the owner of the usage record.
IPAddress *string `json:"ip_address,omitempty"`
// SessionID is the explicit client-provided request correlation identifier
// (e.g. the session_id / X-Session-Id headers). Omitted when absent.
SessionID *string `json:"session_id,omitempty"`
// Cache TTL Override 标记
CacheTTLOverridden bool `json:"cache_ttl_overridden"`
@@ -535,6 +535,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
// ForceCacheBilling 提前拍成标量,避免 worker 闭包保活 failover 状态里的响应体。
forceCacheBilling := fs.ForceCacheBilling
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
sessionID := service.ExtractClientSessionID(c)
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
Result: result,
@@ -547,6 +548,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
UpstreamEndpoint: upstreamEndpoint,
UserAgent: userAgent,
IPAddress: clientIP,
SessionID: sessionID,
RequestPayloadHash: requestPayloadHash,
ForceCacheBilling: forceCacheBilling,
APIKeyService: h.apiKeyService,
@@ -970,6 +972,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
// ForceCacheBilling 提前拍成标量,避免 worker 闭包保活 failover 状态里的响应体。
forceCacheBilling := fs.ForceCacheBilling
quotaPlatform := service.QuotaPlatform(c.Request.Context(), currentAPIKey)
sessionID := service.ExtractClientSessionID(c)
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
Result: result,
@@ -982,6 +985,7 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
UpstreamEndpoint: upstreamEndpoint,
UserAgent: userAgent,
IPAddress: clientIP,
SessionID: sessionID,
RequestPayloadHash: requestPayloadHash,
ForceCacheBilling: forceCacheBilling,
APIKeyService: h.apiKeyService,
@@ -298,6 +298,7 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) {
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
sessionID := service.ExtractClientSessionID(c)
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
Result: result,
@@ -312,6 +313,7 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) {
IPAddress: clientIP,
RequestPayloadHash: requestPayloadHash,
APIKeyService: h.apiKeyService,
SessionID: sessionID,
ChannelUsageFields: clientRequestedUsageFields(c, channelMapping, reqModel, result.UpstreamModel),
}); err != nil {
reqLog.Error("gateway.cc.record_usage_failed",
@@ -276,6 +276,7 @@ func (h *GatewayHandler) Responses(c *gin.Context) {
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
sessionID := service.ExtractClientSessionID(c)
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.RecordUsageInput{
Result: result,
@@ -290,6 +291,7 @@ func (h *GatewayHandler) Responses(c *gin.Context) {
IPAddress: clientIP,
RequestPayloadHash: requestPayloadHash,
APIKeyService: h.apiKeyService,
SessionID: sessionID,
ChannelUsageFields: clientRequestedUsageFields(c, channelMapping, reqModel, result.UpstreamModel),
}); err != nil {
reqLog.Error("gateway.responses.record_usage_failed",
@@ -532,6 +532,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
// ForceCacheBilling 提前拍成标量,避免 worker 闭包保活 failover 状态里的响应体。
forceCacheBilling := fs.ForceCacheBilling
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
sessionID := service.ExtractClientSessionID(c)
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsageWithLongContext(ctx, &service.RecordUsageLongContextInput{
Result: result,
@@ -549,6 +550,7 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
LongContextMultiplier: 2.0, // 超出部分双倍计费
ForceCacheBilling: forceCacheBilling,
APIKeyService: h.apiKeyService,
SessionID: sessionID,
ChannelUsageFields: clientRequestedUsageFields(c, channelMapping, reqModel, result.UpstreamModel),
}); err != nil {
logger.L().With(
+2
View File
@@ -465,6 +465,7 @@ func recordGrokMediaUsage(
) {
userAgent := c.GetHeader("User-Agent")
clientIP := ip.GetClientIP(c)
sessionID := service.ExtractClientSessionID(c)
payloadForHash := body
if len(payloadForHash) == 0 && strings.TrimSpace(requestID) != "" {
payloadForHash = []byte(requestID)
@@ -493,6 +494,7 @@ func recordGrokMediaUsage(
RequestPayloadHash: service.HashUsageRequestPayload(payloadForHash),
APIKeyService: h.apiKeyService,
QuotaPlatform: quotaPlatform,
SessionID: sessionID,
ChannelUsageFields: channelUsageFields,
}); err != nil {
logger.L().With(
@@ -233,6 +233,7 @@ func (h *OpenAIGatewayHandler) recordAlphaSearchUsage(
) {
userAgent := c.GetHeader("User-Agent")
clientIP := ip.GetClientIP(c)
sessionID := service.ExtractClientSessionID(c)
requestPayloadHash := service.HashUsageRequestPayload(body)
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
@@ -252,6 +253,7 @@ func (h *OpenAIGatewayHandler) recordAlphaSearchUsage(
RequestPayloadHash: requestPayloadHash,
APIKeyService: h.apiKeyService,
QuotaPlatform: quotaPlatform,
SessionID: sessionID,
ChannelUsageFields: channelMapping.ToUsageFields(requestedModel, result.UpstreamModel),
}); err != nil {
logger.L().With(
@@ -340,6 +340,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result)
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
sessionID := service.ExtractClientSessionID(c)
cyberBlocked := service.GetOpsCyberPolicy(c) != nil
h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) {
@@ -355,6 +356,7 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
IPAddress: clientIP,
APIKeyService: h.apiKeyService,
QuotaPlatform: quotaPlatform,
SessionID: sessionID,
ChannelUsageFields: clientRequestedUsageFields(c, channelMapping, reqModel, result.UpstreamModel),
CyberBlocked: cyberBlocked,
}); err != nil {
@@ -243,6 +243,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) {
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := GetUpstreamEndpoint(c, account.Platform)
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
sessionID := service.ExtractClientSessionID(c)
h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{
@@ -257,6 +258,7 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) {
IPAddress: clientIP,
APIKeyService: h.apiKeyService,
QuotaPlatform: quotaPlatform,
SessionID: sessionID,
ChannelUsageFields: clientRequestedUsageFields(c, channelMapping, reqModel, result.UpstreamModel),
}); err != nil {
logger.L().With(
@@ -627,6 +627,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result)
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
sessionID := service.ExtractClientSessionID(c)
// 使用量记录通过有界 worker 池提交,避免请求热路径创建无界 goroutine。
cyberBlocked := service.GetOpsCyberPolicy(c) != nil
@@ -644,6 +645,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
RequestPayloadHash: requestPayloadHash,
APIKeyService: h.apiKeyService,
QuotaPlatform: quotaPlatform,
SessionID: sessionID,
ChannelUsageFields: clientRequestedUsageFields(c, channelMapping, reqModel, result.UpstreamModel),
CyberBlocked: cyberBlocked,
}); err != nil {
@@ -1138,6 +1140,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result)
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
sessionID := service.ExtractClientSessionID(c)
cyberBlocked := service.GetOpsCyberPolicy(c) != nil
h.submitOpenAIUsageRecordTask(c.Request.Context(), result, func(ctx context.Context) {
@@ -1154,6 +1157,7 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
RequestPayloadHash: requestPayloadHash,
APIKeyService: h.apiKeyService,
QuotaPlatform: quotaPlatform,
SessionID: sessionID,
ChannelUsageFields: clientRequestedUsageFields(c, channelMappingMsg, reqModel, result.UpstreamModel),
CyberBlocked: cyberBlocked,
}); err != nil {
@@ -1862,6 +1866,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
inboundEndpoint := GetInboundEndpoint(c)
upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result)
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
sessionID := service.ExtractClientSessionID(c)
cyberBlocked := service.GetOpsCyberPolicy(c) != nil
h.submitOpenAIUsageRecordTask(ctx, result, func(taskCtx context.Context) {
if err := h.gatewayService.RecordUsage(taskCtx, &service.OpenAIRecordUsageInput{
@@ -1877,6 +1882,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
RequestPayloadHash: requestPayloadHash,
APIKeyService: h.apiKeyService,
QuotaPlatform: quotaPlatform,
SessionID: sessionID,
ChannelUsageFields: clientRequestedUsageFields(c, channelMappingWS, reqModel, result.UpstreamModel),
CyberBlocked: cyberBlocked,
}); err != nil {
@@ -2884,6 +2890,8 @@ func (h *OpenAIGatewayHandler) recordCyberPolicyIfMarked(c *gin.Context, apiKey
userAgent = c.GetHeader("User-Agent")
clientIPStr = strings.TrimSpace(ip.GetClientIP(c))
}
// 提前拍成标量,避免在下方 goroutine 内访问 gin.Context。
sessionID := service.ExtractClientSessionID(c)
apiKeyPrefix := ""
if apiKey != nil {
apiKeyPrefix = keyPrefix(apiKey.Key, 8)
@@ -2940,6 +2948,7 @@ func (h *OpenAIGatewayHandler) recordCyberPolicyIfMarked(c *gin.Context, apiKey
UpstreamEndpoint: upstreamEndpoint,
UserAgent: userAgent,
IPAddress: clientIPStr,
SessionID: sessionID,
RequestPayloadHash: requestPayloadHash,
APIKeyService: apiKeySvc,
ChannelUsageFields: channelFields,
@@ -374,6 +374,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
if result != nil {
upstreamModel = result.UpstreamModel
}
sessionID := service.ExtractClientSessionID(c)
h.submitMandatoryUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsage(ctx, &service.OpenAIRecordUsageInput{
Result: result,
@@ -388,6 +389,7 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
RequestPayloadHash: requestPayloadHash,
APIKeyService: h.apiKeyService,
QuotaPlatform: quotaPlatform,
SessionID: sessionID,
ChannelUsageFields: clientRequestedUsageFields(c, channelMapping, requestModel, upstreamModel),
}); err != nil {
logger.L().With(
@@ -750,7 +750,7 @@ INSERT INTO batch_image_jobs (
batch_discount_multiplier, hold_multiplier, billable_unit_price, hold_unit_price,
pricing_snapshot_version,
currency, hold_id,
idempotency_key, request_hash, manifest_hash, retry_count, output_expires_at
idempotency_key, request_hash, manifest_hash, retry_count, session_id, output_expires_at
) VALUES (
$1, $2, $3, $4, $5, $6, $7, $8, $9,
$10, $11, $12, $13, $14,
@@ -760,7 +760,7 @@ INSERT INTO batch_image_jobs (
$25, $26, $27, $28,
$29,
$30, $31,
$32, $33, $34, $35, $36
$32, $33, $34, $35, $36, $37
)
RETURNING `+batchImageJobColumns,
params.BatchID, params.UserID, params.APIKeyID, params.AccountID, params.Provider, params.Model, params.TaskName, params.ParentBatchID, params.Status,
@@ -771,7 +771,7 @@ RETURNING `+batchImageJobColumns,
params.BatchDiscountMultiplier, params.HoldMultiplier, params.BillableUnitPrice, params.HoldUnitPrice,
params.PricingSnapshotVersion,
params.Currency, params.HoldID,
params.IdempotencyKey, params.RequestHash, params.ManifestHash, params.RetryCount, params.OutputExpiresAt,
params.IdempotencyKey, params.RequestHash, params.ManifestHash, params.RetryCount, params.SessionID, params.OutputExpiresAt,
))
}
@@ -825,7 +825,7 @@ batch_discount_multiplier, hold_multiplier, billable_unit_price, hold_unit_price
pricing_snapshot_version,
currency, hold_id,
idempotency_key, request_hash, manifest_hash,
retry_count, version, output_expires_at, input_deleted_at, output_deleted_at, downloaded_at, user_deleted_at,
retry_count, version, session_id, output_expires_at, input_deleted_at, output_deleted_at, downloaded_at, user_deleted_at,
last_error_code, last_error_message,
created_at, updated_at, submitted_at, started_at, finished_at, settled_at`
@@ -838,6 +838,7 @@ func scanBatchImageJob(row rowScanner) (*service.BatchImageJob, error) {
var parentBatchID sql.NullString
var holdAmount, actualCost sql.NullFloat64
var holdID, idempotencyKey, requestHash, manifestHash sql.NullString
var sessionID sql.NullString
var outputExpiresAt, inputDeletedAt, outputDeletedAt, downloadedAt, userDeletedAt sql.NullTime
var lastErrorCode, lastErrorMessage sql.NullString
var submittedAt, startedAt, finishedAt, settledAt sql.NullTime
@@ -852,7 +853,7 @@ func scanBatchImageJob(row rowScanner) (*service.BatchImageJob, error) {
&job.PricingSnapshotVersion,
&job.Currency, &holdID,
&idempotencyKey, &requestHash, &manifestHash,
&job.RetryCount, &job.Version, &outputExpiresAt, &inputDeletedAt, &outputDeletedAt, &downloadedAt, &userDeletedAt,
&job.RetryCount, &job.Version, &sessionID, &outputExpiresAt, &inputDeletedAt, &outputDeletedAt, &downloadedAt, &userDeletedAt,
&lastErrorCode, &lastErrorMessage,
&job.CreatedAt, &job.UpdatedAt, &submittedAt, &startedAt, &finishedAt, &settledAt,
)
@@ -874,6 +875,7 @@ func scanBatchImageJob(row rowScanner) (*service.BatchImageJob, error) {
job.IdempotencyKey = batchImageNullStringPtr(idempotencyKey)
job.RequestHash = batchImageNullStringPtr(requestHash)
job.ManifestHash = batchImageNullStringPtr(manifestHash)
job.SessionID = batchImageNullStringPtr(sessionID)
job.OutputExpiresAt = batchImageNullTimePtr(outputExpiresAt)
job.InputDeletedAt = batchImageNullTimePtr(inputDeletedAt)
job.OutputDeletedAt = batchImageNullTimePtr(outputDeletedAt)
@@ -79,6 +79,7 @@ var usageLogInsertArgTypes = [...]string{
"text", // billing_tier
"text", // billing_mode
"numeric", // account_stats_cost
"text", // session_id
"timestamptz", // created_at
}
@@ -274,6 +275,7 @@ func (r *usageLogRepository) createSingle(ctx context.Context, sqlq sqlExecutor,
billing_tier,
billing_mode,
account_stats_cost,
session_id,
created_at
) VALUES (
$1, $2, $3, $4, $5, $6, $7,
@@ -281,7 +283,7 @@ func (r *usageLogRepository) createSingle(ctx context.Context, sqlq sqlExecutor,
$10, $11, $12, $13,
$14, $15, $16, $17,
$18, $19, $20, $21, $22, $23,
$24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53, $54, $55, $56
$24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53, $54, $55, $56, $57
)
ON CONFLICT (request_id, api_key_id) DO NOTHING
RETURNING id, created_at
@@ -728,10 +730,13 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage
billing_tier,
billing_mode,
account_stats_cost,
session_id,
created_at
) AS (VALUES `)
args := make([]any, 0, len(keys)*56)
// Each batch row prepends the synthetic input_index before the 57
// usage-log column values.
args := make([]any, 0, len(keys)*58)
argPos := 1
for idx, key := range keys {
if idx > 0 {
@@ -815,6 +820,7 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage
billing_tier,
billing_mode,
account_stats_cost,
session_id,
created_at
)
SELECT
@@ -873,6 +879,7 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage
billing_tier,
billing_mode,
account_stats_cost,
session_id,
created_at
FROM input
ON CONFLICT (request_id, api_key_id) DO NOTHING
@@ -971,10 +978,11 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) (
billing_tier,
billing_mode,
account_stats_cost,
session_id,
created_at
) AS (VALUES `)
args := make([]any, 0, len(preparedList)*56)
args := make([]any, 0, len(preparedList)*57)
argPos := 1
for idx, prepared := range preparedList {
if idx > 0 {
@@ -1055,6 +1063,7 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) (
billing_tier,
billing_mode,
account_stats_cost,
session_id,
created_at
)
SELECT
@@ -1113,6 +1122,7 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) (
billing_tier,
billing_mode,
account_stats_cost,
session_id,
created_at
FROM input
ON CONFLICT (request_id, api_key_id) DO NOTHING
@@ -1179,6 +1189,7 @@ func execUsageLogInsertNoResult(ctx context.Context, sqlq sqlExecutor, prepared
billing_tier,
billing_mode,
account_stats_cost,
session_id,
created_at
) VALUES (
$1, $2, $3, $4, $5, $6, $7,
@@ -1186,7 +1197,7 @@ func execUsageLogInsertNoResult(ctx context.Context, sqlq sqlExecutor, prepared
$10, $11, $12, $13,
$14, $15, $16, $17,
$18, $19, $20, $21, $22, $23,
$24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53, $54, $55, $56
$24, $25, $26, $27, $28, $29, $30, $31, $32, $33, $34, $35, $36, $37, $38, $39, $40, $41, $42, $43, $44, $45, $46, $47, $48, $49, $50, $51, $52, $53, $54, $55, $56, $57
)
ON CONFLICT (request_id, api_key_id) DO NOTHING
`, prepared.args...)
@@ -1227,6 +1238,7 @@ func prepareUsageLogInsert(log *service.UsageLog) usageLogInsertPrepared {
modelMappingChain := nullString(log.ModelMappingChain)
billingTier := nullString(log.BillingTier)
billingMode := nullString(log.BillingMode)
sessionID := nullString(log.SessionID)
requestedModel := strings.TrimSpace(log.RequestedModel)
if requestedModel == "" {
requestedModel = strings.TrimSpace(log.Model)
@@ -1299,6 +1311,7 @@ func prepareUsageLogInsert(log *service.UsageLog) usageLogInsertPrepared {
billingTier,
billingMode,
log.AccountStatsCost, // account_stats_cost
sessionID, // session_id
createdAt,
},
}
@@ -19,7 +19,7 @@ import (
"github.com/Wei-Shaw/sub2api/internal/service"
)
const usageLogSelectColumns = "id, user_id, api_key_id, account_id, request_id, model, requested_model, upstream_model, group_id, subscription_id, input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens, cache_creation_5m_tokens, cache_creation_1h_tokens, image_output_tokens, image_output_cost, image_input_tokens, image_input_cost, input_cost, output_cost, cache_creation_cost, cache_read_cost, total_cost, actual_cost, rate_multiplier, account_rate_multiplier, billing_type, request_type, stream, openai_ws_mode, duration_ms, first_token_ms, user_agent, ip_address, image_count, image_size, image_input_size, image_output_size, image_size_source, image_size_breakdown, video_count, video_resolution, video_duration_seconds, service_tier, reasoning_effort, inbound_endpoint, upstream_endpoint, cache_ttl_overridden, long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, billing_mode, account_stats_cost, created_at"
const usageLogSelectColumns = "id, user_id, api_key_id, account_id, request_id, model, requested_model, upstream_model, group_id, subscription_id, input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens, cache_creation_5m_tokens, cache_creation_1h_tokens, image_output_tokens, image_output_cost, image_input_tokens, image_input_cost, input_cost, output_cost, cache_creation_cost, cache_read_cost, total_cost, actual_cost, rate_multiplier, account_rate_multiplier, billing_type, request_type, stream, openai_ws_mode, duration_ms, first_token_ms, user_agent, ip_address, image_count, image_size, image_input_size, image_output_size, image_size_source, image_size_breakdown, video_count, video_resolution, video_duration_seconds, service_tier, reasoning_effort, inbound_endpoint, upstream_endpoint, cache_ttl_overridden, long_context_billing_applied, channel_id, model_mapping_chain, billing_tier, billing_mode, account_stats_cost, session_id, created_at"
func (r *usageLogRepository) GetByID(ctx context.Context, id int64) (log *service.UsageLog, err error) {
query := "SELECT " + usageLogSelectColumns + " FROM usage_logs WHERE id = $1"
@@ -481,6 +481,7 @@ func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, e
billingTier sql.NullString
billingMode sql.NullString
accountStatsCost sql.NullFloat64
sessionID sql.NullString
createdAt time.Time
)
@@ -541,6 +542,7 @@ func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, e
&billingTier,
&billingMode,
&accountStatsCost,
&sessionID,
&createdAt,
); err != nil {
return nil, err
@@ -661,6 +663,9 @@ func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, e
if accountStatsCost.Valid {
log.AccountStatsCost = &accountStatsCost.Float64
}
if sessionID.Valid {
log.SessionID = &sessionID.String
}
return log, nil
}
@@ -96,6 +96,7 @@ func TestUsageLogRepositoryCreateSyncRequestTypeAndLegacyFields(t *testing.T) {
sqlmock.AnyArg(), // billing_tier
sqlmock.AnyArg(), // billing_mode
sqlmock.AnyArg(), // account_stats_cost
sqlmock.AnyArg(), // session_id
createdAt,
).
WillReturnRows(sqlmock.NewRows([]string{"id", "created_at"}).AddRow(int64(99), createdAt))
@@ -185,6 +186,7 @@ func TestUsageLogRepositoryCreate_PersistsServiceTier(t *testing.T) {
sqlmock.AnyArg(), // billing_tier
sqlmock.AnyArg(), // billing_mode
sqlmock.AnyArg(), // account_stats_cost
sqlmock.AnyArg(), // session_id
createdAt,
).
WillReturnRows(sqlmock.NewRows([]string{"id", "created_at"}).AddRow(int64(100), createdAt))
@@ -826,6 +828,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) {
sql.NullString{},
sql.NullString{},
sql.NullFloat64{},
sql.NullString{},
now,
}})
require.NoError(t, err)
@@ -900,6 +903,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) {
sql.NullString{}, // billing_tier
sql.NullString{}, // billing_mode
sql.NullFloat64{}, // account_stats_cost
sql.NullString{}, // session_id
now,
}})
require.NoError(t, err)
@@ -957,6 +961,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) {
sql.NullString{}, // billing_tier
sql.NullString{}, // billing_mode
sql.NullFloat64{}, // account_stats_cost
sql.NullString{}, // session_id
now,
}})
require.NoError(t, err)
@@ -1014,6 +1019,7 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) {
sql.NullString{}, // billing_tier
sql.NullString{}, // billing_mode
sql.NullFloat64{}, // account_stats_cost
sql.NullString{}, // session_id
now,
}})
require.NoError(t, err)
@@ -0,0 +1,71 @@
//go:build integration
package repository
import (
"context"
"testing"
"time"
"github.com/google/uuid"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/stretchr/testify/require"
)
// TestUsageLog_SessionIDPersistence proves session_id round-trips from insert to
// read and is omitted (NULL) when absent.
func TestUsageLog_SessionIDPersistence(t *testing.T) {
ctx := context.Background()
client := testEntClient(t)
repo := newUsageLogRepositoryWithSQL(client, integrationDB)
user := mustCreateUser(t, client, &service.User{Email: "session-id-" + uuid.NewString() + "@example.com"})
apiKey := mustCreateApiKey(t, client, &service.APIKey{UserID: user.ID, Key: "sk-session-" + uuid.NewString(), Name: "k"})
account := mustCreateAccount(t, client, &service.Account{Name: "acc-session-" + uuid.NewString()})
sessionID := "sess-" + uuid.NewString()
withSession := &service.UsageLog{
UserID: user.ID,
APIKeyID: apiKey.ID,
AccountID: account.ID,
RequestID: uuid.NewString(),
Model: "claude-3",
InputTokens: 10,
OutputTokens: 5,
TotalCost: 1.0,
ActualCost: 1.0,
SessionID: &sessionID,
CreatedAt: time.Now().UTC(),
}
_, err := repo.Create(ctx, withSession)
require.NoError(t, err)
require.NotZero(t, withSession.ID)
withoutSession := &service.UsageLog{
UserID: user.ID,
APIKeyID: apiKey.ID,
AccountID: account.ID,
RequestID: uuid.NewString(),
Model: "claude-3",
InputTokens: 7,
OutputTokens: 3,
TotalCost: 0.5,
ActualCost: 0.5,
CreatedAt: time.Now().UTC(),
}
_, err = repo.Create(ctx, withoutSession)
require.NoError(t, err)
// Round-trip: session id survives insert → read.
got, err := repo.GetByID(ctx, withSession.ID)
require.NoError(t, err)
require.NotNil(t, got.SessionID)
require.Equal(t, sessionID, *got.SessionID)
// Omission: absent session id reads back as nil (NULL), not empty string.
gotNone, err := repo.GetByID(ctx, withoutSession.ID)
require.NoError(t, err)
require.Nil(t, gotNone.SessionID)
}
@@ -0,0 +1,91 @@
//go:build unit
package repository
import (
"database/sql"
"strings"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/stretchr/testify/require"
)
func newSessionIDUsageLog(sessionID *string) *service.UsageLog {
return &service.UsageLog{
UserID: 1,
APIKeyID: 2,
AccountID: 3,
RequestID: "req-session-id",
Model: "claude-3",
InputTokens: 10,
OutputTokens: 5,
TotalCost: 1.0,
ActualCost: 1.0,
SessionID: sessionID,
CreatedAt: time.Now().UTC(),
}
}
// TestPrepareUsageLogInsert_SessionIDArgWiring pins the session_id column to the
// arg slice / arg-type table so the five INSERT column lists stay in sync. session_id
// is the penultimate arg (created_at is always last).
func TestPrepareUsageLogInsert_SessionIDArgWiring(t *testing.T) {
require.Len(t, usageLogInsertArgTypes, 57, "arg-type table must include session_id")
sessionID := "sess-persisted-123"
prepared := prepareUsageLogInsert(newSessionIDUsageLog(&sessionID))
require.Len(t, prepared.args, len(usageLogInsertArgTypes),
"prepared args must match the arg-type table length")
// created_at is last; session_id is the arg immediately before it.
sessionArg := prepared.args[len(prepared.args)-2]
ns, ok := sessionArg.(sql.NullString)
require.True(t, ok, "session_id arg should be a sql.NullString, got %T", sessionArg)
require.True(t, ns.Valid)
require.Equal(t, sessionID, ns.String)
require.Equal(t, "text", usageLogInsertArgTypes[len(usageLogInsertArgTypes)-2],
"session_id arg type must be text")
}
// TestPrepareUsageLogInsert_SessionIDNullWhenAbsent proves an absent session id is
// persisted as SQL NULL rather than an empty string.
func TestPrepareUsageLogInsert_SessionIDNullWhenAbsent(t *testing.T) {
prepared := prepareUsageLogInsert(newSessionIDUsageLog(nil))
sessionArg := prepared.args[len(prepared.args)-2]
ns, ok := sessionArg.(sql.NullString)
require.True(t, ok, "session_id arg should be a sql.NullString, got %T", sessionArg)
require.False(t, ns.Valid, "absent session id must be NULL, not empty string")
empty := ""
preparedEmpty := prepareUsageLogInsert(newSessionIDUsageLog(&empty))
nsEmpty := preparedEmpty.args[len(preparedEmpty.args)-2].(sql.NullString)
require.False(t, nsEmpty.Valid, "empty session id must also be NULL")
}
// TestUsageLogInsertQueries_IncludeSessionID guards that every generated INSERT path
// and the SELECT column list reference session_id.
func TestUsageLogInsertQueries_IncludeSessionID(t *testing.T) {
require.Contains(t, usageLogSelectColumns, "session_id",
"SELECT column list must include session_id")
sessionID := "sess-in-query"
log := newSessionIDUsageLog(&sessionID)
prepared := prepareUsageLogInsert(log)
key := usageLogBatchKey(log.RequestID, log.APIKeyID)
batchQuery, batchArgs := buildUsageLogBatchInsertQuery([]string{key},
map[string]usageLogInsertPrepared{key: prepared})
require.Contains(t, batchQuery, "session_id")
// Two column references (INSERT column list + SELECT ... FROM input) plus the CTE def.
require.GreaterOrEqual(t, strings.Count(batchQuery, "session_id"), 3)
require.Len(t, batchArgs, len(prepared.args)+1,
"batch args include the synthetic input_index before usage-log values")
bestEffortQuery, bestEffortArgs := buildUsageLogBestEffortInsertQuery([]usageLogInsertPrepared{prepared})
require.Contains(t, bestEffortQuery, "session_id")
require.Len(t, bestEffortArgs, len(prepared.args))
}
+2
View File
@@ -137,6 +137,7 @@ type BatchImageJob struct {
IdempotencyKey *string
RequestHash *string
ManifestHash *string
SessionID *string
RetryCount int
Version int
@@ -196,6 +197,7 @@ type CreateBatchImageJobParams struct {
IdempotencyKey *string
RequestHash *string
ManifestHash *string
SessionID *string
RetryCount int
@@ -425,6 +425,7 @@ func (r *fakeBatchImageRepository) CreateBatchImageJob(_ context.Context, params
Currency: params.Currency,
IdempotencyKey: params.IdempotencyKey,
RequestHash: params.RequestHash,
SessionID: params.SessionID,
CreatedAt: time.Now(),
}
r.jobs[job.BatchID] = job
@@ -57,6 +57,7 @@ type BatchImageSubmitRequest struct {
AspectRatio string `json:"aspect_ratio"`
ImageSize string `json:"image_size"`
Metadata map[string]string `json:"metadata"`
SessionID *string `json:"-"`
}
type BatchImageSubmitItem struct {
@@ -281,6 +282,7 @@ func (s *BatchImagePublicService) Submit(ctx context.Context, owner BatchImageOw
HoldID: &holdID,
IdempotencyKey: batchImageOptionalStringPtr(idempotencyKey),
RequestHash: batchImageStringPtr(requestHash),
SessionID: normalized.SessionID,
})
if err != nil {
return nil, err
@@ -27,8 +27,10 @@ func TestBatchImagePublicService_Submit(t *testing.T) {
t.Run("accepts valid request stores refs and enqueues once", func(t *testing.T) {
svc, repo, queue, gemini, _ := newTestBatchImagePublicService(true)
req := validBatchImageSubmitRequest()
req.SessionID = batchImageStringPtr("batch-session-123")
got, err := svc.Submit(ctx, testBatchImageOwner(), validBatchImageSubmitRequest(), "")
got, err := svc.Submit(ctx, testBatchImageOwner(), req, "")
require.NoError(t, err)
require.Equal(t, "image.batch", got.Object)
require.Equal(t, "queued", got.Status)
@@ -61,6 +63,7 @@ func TestBatchImagePublicService_Submit(t *testing.T) {
require.InDelta(t, 0.6, job.HoldMultiplier, 1e-12)
require.InDelta(t, 0.125, job.BillableUnitPrice, 1e-12)
require.InDelta(t, 0.15, job.HoldUnitPrice, 1e-12)
require.Equal(t, "batch-session-123", batchImageDerefString(job.SessionID))
})
t.Run("combines user group image rate account rate discount and hold margin", func(t *testing.T) {
@@ -389,15 +392,18 @@ func TestBatchImagePublicService_Submit(t *testing.T) {
})
t.Run("idempotency returns same batch without provider resubmit", func(t *testing.T) {
svc, _, queue, gemini, _ := newTestBatchImagePublicService(true)
svc, repo, queue, gemini, _ := newTestBatchImagePublicService(true)
req := validBatchImageSubmitRequest()
req.SessionID = batchImageStringPtr("original-session")
first, err := svc.Submit(ctx, testBatchImageOwner(), req, "client-key")
require.NoError(t, err)
req.SessionID = batchImageStringPtr("retry-session")
second, err := svc.Submit(ctx, testBatchImageOwner(), req, "client-key")
require.NoError(t, err)
require.Equal(t, first.ID, second.ID)
require.Equal(t, "original-session", batchImageDerefString(repo.jobs[first.ID].SessionID))
require.Len(t, gemini.submits, 1)
require.Equal(t, []string{first.ID}, queue.enqueued)
})
@@ -277,6 +277,7 @@ func (s *BatchImageSettlementService) recordUsageLog(ctx context.Context, job *B
RequestType: RequestTypeSync,
BillingMode: &billingMode,
ImageSize: &imageSize,
SessionID: job.SessionID,
CreatedAt: createdAt,
}
writeUsageLogBestEffort(ctx, s.UsageLogRepo, usageLog, "service.batch_image_settlement")
@@ -18,9 +18,14 @@ func TestBatchImageSettlementService_SettlesAndChargesSuccessfulImagesOnly(t *te
job.SuccessCount = 3
job.FailCount = 2
job.ItemCount = 5
job.SessionID = batchImageStringPtr("batch-settlement-session")
repo.jobs[job.BatchID] = job
billing := &fakeBatchImageBillingRepo{}
svc := &BatchImageSettlementService{Repo: repo, BillingRepo: billing, Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.25}}
usageLogs := &openAIRecordUsageLogRepoStub{}
svc := &BatchImageSettlementService{
Repo: repo, BillingRepo: billing, Pricing: &fakeBatchImagePricingResolver{unitPrice: 0.25},
UsageLogRepo: usageLogs,
}
result, err := svc.Settle(context.Background(), job.BatchID)
require.NoError(t, err)
@@ -32,6 +37,7 @@ func TestBatchImageSettlementService_SettlesAndChargesSuccessfulImagesOnly(t *te
require.Equal(t, 0.75, *repo.jobs[job.BatchID].ActualCost)
require.NotEmpty(t, batchImageDerefString(repo.jobs[job.BatchID].ManifestHash))
require.NotNil(t, repo.jobs[job.BatchID].SettledAt)
require.Equal(t, "batch-settlement-session", batchImageDerefString(usageLogs.lastLog.SessionID))
require.Len(t, billing.captures, 1)
require.Equal(t, int64(321), billing.captures[0].APIKeyID)
require.Equal(t, job.UserID, billing.captures[0].UserID)
@@ -46,6 +46,7 @@ type RecordUsageInput struct {
UpstreamEndpoint string // 上游端点(标准化后的上游路径)
UserAgent string // 请求的 User-Agent
IPAddress string // 请求的客户端 IP 地址
SessionID string // 客户端显式会话标识(session_id / X-Session-Id 等请求头),仅用于用量行会话关联
RequestPayloadHash string // 请求体语义哈希,用于降低 request_id 误复用时的静默误去重风险
ForceCacheBilling bool // 强制缓存计费:将 input_tokens 转为 cache_read 计费(用于粘性会话切换)
APIKeyService APIKeyQuotaUpdater // 可选:用于更新API Key配额
@@ -572,6 +573,7 @@ func (s *GatewayService) RecordUsage(ctx context.Context, input *RecordUsageInpu
UpstreamEndpoint: input.UpstreamEndpoint,
UserAgent: input.UserAgent,
IPAddress: input.IPAddress,
SessionID: input.SessionID,
RequestPayloadHash: input.RequestPayloadHash,
ForceCacheBilling: input.ForceCacheBilling,
APIKeyService: input.APIKeyService,
@@ -591,6 +593,7 @@ type RecordUsageLongContextInput struct {
UpstreamEndpoint string // 上游端点(标准化后的上游路径)
UserAgent string // 请求的 User-Agent
IPAddress string // 请求的客户端 IP 地址
SessionID string // 客户端显式会话标识(session_id / X-Session-Id 等请求头),仅用于用量行会话关联
RequestPayloadHash string // 请求体语义哈希,用于降低 request_id 误复用时的静默误去重风险
LongContextThreshold int // 长上下文阈值(如 200000)
LongContextMultiplier float64 // 超出阈值部分的倍率(如 2.0)
@@ -613,6 +616,7 @@ func (s *GatewayService) RecordUsageWithLongContext(ctx context.Context, input *
UpstreamEndpoint: input.UpstreamEndpoint,
UserAgent: input.UserAgent,
IPAddress: input.IPAddress,
SessionID: input.SessionID,
RequestPayloadHash: input.RequestPayloadHash,
ForceCacheBilling: input.ForceCacheBilling,
APIKeyService: input.APIKeyService,
@@ -635,6 +639,7 @@ type recordUsageCoreInput struct {
UpstreamEndpoint string
UserAgent string
IPAddress string
SessionID string
RequestPayloadHash string
ForceCacheBilling bool
APIKeyService APIKeyQuotaUpdater
@@ -1010,6 +1015,7 @@ func (s *GatewayService) buildRecordUsageLog(
ModelMappingChain: optionalTrimmedStringPtr(input.ModelMappingChain),
UserAgent: optionalTrimmedStringPtr(input.UserAgent),
IPAddress: optionalTrimmedStringPtr(input.IPAddress),
SessionID: optionalTrimmedStringPtr(input.SessionID),
GroupID: apiKey.GroupID,
SubscriptionID: optionalSubscriptionID(subscription),
CreatedAt: time.Now(),
@@ -125,6 +125,11 @@ func isGrokRequestContext(c *gin.Context) bool {
if c == nil {
return false
}
if c.Request != nil {
if platform, ok := ResolvedTargetPlatformFromContext(c.Request.Context()); ok {
return platform == PlatformGrok
}
}
v, exists := c.Get("api_key")
if !exists {
return false
@@ -26,6 +26,15 @@ const (
codeBuddyConversationHeader = "X-Conversation-ID"
)
var explicitOpenAIHeaderSessionNames = []string{
"session_id",
"conversation_id",
openCodeSessionAffinityHeader,
openCodeSessionIDHeader,
openCodeNativeSessionHeader,
codeBuddyConversationHeader,
}
// explicitOpenAIHeaderSessionID resolves stable conversation identifiers sent
// by OpenAI-compatible clients. Keep this list limited to session-scoped
// fields: request/message IDs rotate every turn and would defeat sticky routing
@@ -35,14 +44,7 @@ func explicitOpenAIHeaderSessionID(c *gin.Context) string {
return ""
}
for _, header := range []string{
"session_id",
"conversation_id",
openCodeSessionAffinityHeader,
openCodeSessionIDHeader,
openCodeNativeSessionHeader,
codeBuddyConversationHeader,
} {
for _, header := range explicitOpenAIHeaderSessionNames {
if sessionID := strings.TrimSpace(c.GetHeader(header)); sessionID != "" {
return sessionID
}
@@ -29,6 +29,7 @@ type OpenAIRecordUsageInput struct {
UpstreamEndpoint string
UserAgent string // 请求的 User-Agent
IPAddress string // 请求的客户端 IP 地址
SessionID string // 客户端显式会话标识(session_id / X-Session-Id 等请求头),仅用于用量行会话关联
RequestPayloadHash string
APIKeyService APIKeyQuotaUpdater
QuotaPlatform string // user×platform quota platform resolved by the handler before async billing.
@@ -55,6 +56,7 @@ type CyberPolicyUsageInput struct {
UpstreamEndpoint string
UserAgent string
IPAddress string
SessionID string
RequestPayloadHash string
APIKeyService APIKeyQuotaUpdater
ChannelUsageFields
@@ -89,6 +91,7 @@ func (s *OpenAIGatewayService) RecordCyberPolicyUsageLog(ctx context.Context, in
UpstreamEndpoint: in.UpstreamEndpoint,
UserAgent: in.UserAgent,
IPAddress: in.IPAddress,
SessionID: in.SessionID,
RequestPayloadHash: in.RequestPayloadHash,
APIKeyService: in.APIKeyService,
ChannelUsageFields: in.ChannelUsageFields,
@@ -333,6 +336,9 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
usageLog.IPAddress = &input.IPAddress
}
// 添加 SessionID(客户端显式会话标识;缺失/无效时保持 nil)
usageLog.SessionID = optionalTrimmedStringPtr(input.SessionID)
if apiKey.GroupID != nil {
usageLog.GroupID = apiKey.GroupID
}
+75
View File
@@ -0,0 +1,75 @@
package service
import (
"strings"
"unicode/utf8"
"github.com/gin-gonic/gin"
)
// maxPersistedSessionIDLength bounds the persisted client session identifier to the
// usage_logs.session_id column width (VARCHAR(255)). Longer values are rejected so
// distinct identifiers can never alias through truncation.
const maxPersistedSessionIDLength = 255
// clientSessionIDHeaders extends the OpenAI-compatible sticky-session signals with
// native protocol identifiers that are safe to persist but must not alter OpenAI
// scheduling behavior.
var clientSessionIDHeaders = append(
append([]string(nil), explicitOpenAIHeaderSessionNames...),
claudeCodeSessionHeader,
)
// ExtractClientSessionID resolves the explicit client-provided session identifier from
// request headers for usage-log correlation and returns it sanitized. It is
// protocol-agnostic and shared by every gateway handler so all supported protocols
// record session_id through one seam. Returns "" when no valid identifier is present.
//
// This value feeds only usage_logs.session_id persistence. It does NOT affect sticky
// routing, account selection, request_id semantics, or upstream prompt caching, which
// keep their own (intentionally broader) session-signal resolution.
func ExtractClientSessionID(c *gin.Context) string {
if c == nil || c.Request == nil {
return ""
}
for _, header := range clientSessionIDHeaders {
if sessionID := sanitizeSessionID(c.GetHeader(header)); sessionID != "" {
return sessionID
}
}
if isGrokRequestContext(c) {
if sessionID := sanitizeSessionID(c.GetHeader(grokConversationIDHeader)); sessionID != "" {
return sessionID
}
}
return ""
}
// sanitizeSessionID normalizes a raw client-supplied session identifier for safe
// persistence: it trims surrounding whitespace, rejects the value outright if it
// contains any control character (CR/LF/tab/NUL/…) so a log- or header-injection style
// payload cannot slip into stored correlation data, and rejects values longer than
// the DB column bound. Absent or invalid input yields "".
func sanitizeSessionID(raw string) string {
if !utf8.ValidString(raw) {
return ""
}
trimmed := strings.TrimSpace(raw)
if trimmed == "" {
return ""
}
count := 0
for _, r := range trimmed {
if r < 0x20 || r == 0x7f {
// An explicit correlation id never legitimately contains control
// characters; drop the whole value rather than persist a mangled or
// partially-injected identifier.
return ""
}
count++
if count > maxPersistedSessionIDLength {
return ""
}
}
return trimmed
}
+155
View File
@@ -0,0 +1,155 @@
//go:build unit
package service
import (
"context"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
func newSessionHeaderContext(t *testing.T, headers map[string]string) *gin.Context {
t.Helper()
gin.SetMode(gin.TestMode)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
req := httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
for k, v := range headers {
req.Header.Set(k, v)
}
c.Request = req
return c
}
func TestSanitizeSessionID(t *testing.T) {
longRunes := strings.Repeat("a", maxPersistedSessionIDLength+50)
multibyte := strings.Repeat("好", maxPersistedSessionIDLength+10)
tests := []struct {
name string
in string
want string
}{
{"empty", "", ""},
{"whitespace only", " \t ", ""},
{"trims surrounding whitespace", " sess-123 ", "sess-123"},
{"plain value", "conv_abc-123.XYZ", "conv_abc-123.XYZ"},
{"uuid", "550e8400-e29b-41d4-a716-446655440000", "550e8400-e29b-41d4-a716-446655440000"},
{"reject CR", "sess\r123", ""},
{"reject LF", "sess\n123", ""},
{"reject CRLF injection", "sess-1\r\nSet-Cookie: x=y", ""},
{"reject tab inside", "sess\t123", ""},
{"reject NUL", "sess\x00123", ""},
{"reject DEL", "sess\x7f123", ""},
{"reject invalid UTF-8", string([]byte{'s', 'e', 's', 's', '-', 0xff}), ""},
{"accepts value at column bound", strings.Repeat("b", maxPersistedSessionIDLength), strings.Repeat("b", maxPersistedSessionIDLength)},
{"rejects overlong value", longRunes, ""},
{"rejects overlong multibyte value", multibyte, ""},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
got := sanitizeSessionID(tc.in)
require.Equal(t, tc.want, got, "sanitizeSessionID(%q)", tc.in)
// Sanitized output must never exceed the DB column bound (rune-counted).
require.LessOrEqual(t, len([]rune(got)), maxPersistedSessionIDLength)
})
}
}
func TestExtractClientSessionID_NilContext(t *testing.T) {
require.Equal(t, "", ExtractClientSessionID(nil))
}
func TestExtractClientSessionID_NilRequest(t *testing.T) {
require.Equal(t, "", ExtractClientSessionID(&gin.Context{}))
}
func TestExtractClientSessionID_AbsentReturnsEmpty(t *testing.T) {
c := newSessionHeaderContext(t, nil)
require.Equal(t, "", ExtractClientSessionID(c))
}
func TestExtractClientSessionID_SupportedHeaders(t *testing.T) {
tests := []struct {
name string
header string
value string
}{
{"session_id", "session_id", "sess-A"},
{"conversation_id", "conversation_id", "conv-B"},
{"X-Session-Affinity", openCodeSessionAffinityHeader, "aff-C"},
{"X-Session-Id", openCodeSessionIDHeader, "sid-D"},
{"X-OpenCode-Session", openCodeNativeSessionHeader, "oc-E"},
{"X-Conversation-ID", codeBuddyConversationHeader, "cb-F"},
{"X-Claude-Code-Session-Id", claudeCodeSessionHeader, "cc-G"},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
c := newSessionHeaderContext(t, map[string]string{tc.header: tc.value})
require.Equal(t, tc.value, ExtractClientSessionID(c))
})
}
}
func TestExtractClientSessionID_HeaderPrecedence(t *testing.T) {
// session_id ranks ahead of conversation_id and the X-* variants.
c := newSessionHeaderContext(t, map[string]string{
"session_id": "primary",
"conversation_id": "secondary",
openCodeSessionIDHeader: "tertiary",
codeBuddyConversationHeader: "quaternary",
})
require.Equal(t, "primary", ExtractClientSessionID(c))
}
func TestExtractClientSessionID_Sanitizes(t *testing.T) {
c := newSessionHeaderContext(t, map[string]string{openCodeSessionIDHeader: " clean-123 "})
require.Equal(t, "clean-123", ExtractClientSessionID(c))
}
func TestExtractClientSessionID_IgnoresNonSessionHeaders(t *testing.T) {
// prompt_cache_key, request/message ids, and a Grok conversation header on a
// non-Grok request are NOT persisted as session_id.
c := newSessionHeaderContext(t, map[string]string{
"prompt_cache_key": "cache-key-should-not-persist",
"X-Request-Id": "req-should-not-persist",
"x-grok-conv-id": "grok-conv-should-not-persist",
})
require.Equal(t, "", ExtractClientSessionID(c))
}
func TestExtractClientSessionID_GrokConversationHeader(t *testing.T) {
c := newSessionHeaderContext(t, map[string]string{
grokConversationIDHeader: "grok-native-session",
})
c.Set("api_key", &APIKey{
ID: 42,
Group: &Group{Platform: PlatformGrok},
})
require.Equal(t, "grok-native-session", ExtractClientSessionID(c))
}
func TestExtractClientSessionID_GrokConversationHeaderForCompositeRoute(t *testing.T) {
c := newSessionHeaderContext(t, map[string]string{
grokConversationIDHeader: "grok-composite-session",
})
c.Set("api_key", &APIKey{
ID: 43,
Group: &Group{Platform: PlatformComposite},
})
c.Request = c.Request.WithContext(WithResolvedTargetPlatform(context.Background(), PlatformGrok))
require.Equal(t, "grok-composite-session", ExtractClientSessionID(c))
}
func TestExtractClientSessionID_InjectionHeaderDropped(t *testing.T) {
// A supported header carrying a CRLF payload is rejected, not persisted mangled.
c := newSessionHeaderContext(t, map[string]string{"session_id": "abc"})
c.Request.Header.Set("session_id", "abc\r\nX-Injected: 1")
require.Equal(t, "", ExtractClientSessionID(c))
}
+4
View File
@@ -165,6 +165,10 @@ type UsageLog struct {
FirstTokenMs *int
UserAgent *string
IPAddress *string
// SessionID is the explicit client-provided request correlation identifier
// (e.g. the session_id / X-Session-Id headers). Nil when the client sent no
// valid session header. It is never derived from prompt_cache_key or content.
SessionID *string
// Cache TTL Override 标记(管理员强制替换了缓存 TTL 计费)
CacheTTLOverridden bool
@@ -0,0 +1,14 @@
-- Persist the explicit client-provided request session/conversation identifier
-- (e.g. the session_id / X-Session-Id / X-Conversation-ID request headers) so
-- usage rows can be correlated by session for querying.
--
-- Nullable with no default: on PostgreSQL 11+ this is a metadata-only change, so
-- it does NOT rewrite the (potentially large) usage_logs table. Absent/invalid
-- session identifiers stay NULL. Only explicit client correlation identity is
-- stored here; prompt_cache_key and content-derived sticky hashes are never
-- persisted as session_id.
ALTER TABLE usage_logs ADD COLUMN IF NOT EXISTS session_id VARCHAR(255);
-- Batch image usage is recorded asynchronously after the originating HTTP
-- request has finished, so retain the same sanitized identifier on the job.
ALTER TABLE batch_image_jobs ADD COLUMN IF NOT EXISTS session_id VARCHAR(255);