mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 14:38:27 +08:00
feat(usage): persist client session identifiers
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
@@ -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);
|
||||
Reference in New Issue
Block a user