Fix image billing size normalization

This commit is contained in:
2ue
2026-05-12 15:21:31 +08:00
parent 62ccd0ff39
commit bb4c1abe28
45 changed files with 3270 additions and 313 deletions
+4
View File
@@ -600,6 +600,10 @@ func usageLogFromServiceUser(l *service.UsageLog) UsageLog {
FirstTokenMs: l.FirstTokenMs,
ImageCount: l.ImageCount,
ImageSize: l.ImageSize,
ImageInputSize: l.ImageInputSize,
ImageOutputSize: l.ImageOutputSize,
ImageSizeSource: l.ImageSizeSource,
ImageSizeBreakdown: l.ImageSizeBreakdown,
MediaType: l.MediaType,
UserAgent: l.UserAgent,
CacheTTLOverridden: l.CacheTTLOverridden,
@@ -148,6 +148,65 @@ func TestUsageLogFromService_FallsBackToLegacyModelWhenRequestedModelMissing(t *
require.Equal(t, "claude-3", adminDTO.Model)
}
func TestUsageLogFromService_IncludesImageBillingMetadataForUserAndAdmin(t *testing.T) {
t.Parallel()
imageSize := "4K"
inputSize := "1024x1024"
outputSize := "3840x2160"
source := "output"
log := &service.UsageLog{
RequestID: "req_image_metadata",
Model: "gpt-image-2",
ImageCount: 2,
ImageSize: &imageSize,
ImageInputSize: &inputSize,
ImageOutputSize: &outputSize,
ImageSizeSource: &source,
ImageSizeBreakdown: map[string]int{"4K": 2},
}
userDTO := UsageLogFromService(log)
adminDTO := UsageLogFromServiceAdmin(log)
for _, got := range []*UsageLog{userDTO, &adminDTO.UsageLog} {
require.Equal(t, 2, got.ImageCount)
require.NotNil(t, got.ImageSize)
require.Equal(t, imageSize, *got.ImageSize)
require.NotNil(t, got.ImageInputSize)
require.Equal(t, inputSize, *got.ImageInputSize)
require.NotNil(t, got.ImageOutputSize)
require.Equal(t, outputSize, *got.ImageOutputSize)
require.NotNil(t, got.ImageSizeSource)
require.Equal(t, source, *got.ImageSizeSource)
require.Equal(t, map[string]int{"4K": 2}, got.ImageSizeBreakdown)
}
}
func TestUsageLogFromService_PreservesHistoricalMissingImageSize(t *testing.T) {
t.Parallel()
log := &service.UsageLog{
RequestID: "req_legacy_image_missing_size",
Model: "gpt-image-2",
ImageCount: 1,
ImageSize: nil,
}
dto := UsageLogFromService(log)
require.Equal(t, 1, dto.ImageCount)
require.Nil(t, dto.ImageSize)
require.Nil(t, dto.ImageInputSize)
require.Nil(t, dto.ImageOutputSize)
require.Nil(t, dto.ImageSizeSource)
require.Nil(t, dto.ImageSizeBreakdown)
body, err := json.Marshal(dto)
require.NoError(t, err)
require.Contains(t, string(body), `"image_size":null`)
require.NotContains(t, string(body), `"image_size":"2K"`)
}
func f64Ptr(value float64) *float64 {
return &value
}
+7 -3
View File
@@ -400,9 +400,13 @@ type UsageLog struct {
FirstTokenMs *int `json:"first_token_ms"`
// 图片生成字段
ImageCount int `json:"image_count"`
ImageSize *string `json:"image_size"`
MediaType *string `json:"media_type"`
ImageCount int `json:"image_count"`
ImageSize *string `json:"image_size"`
ImageInputSize *string `json:"image_input_size"`
ImageOutputSize *string `json:"image_output_size"`
ImageSizeSource *string `json:"image_size_source"`
ImageSizeBreakdown map[string]int `json:"image_size_breakdown"`
MediaType *string `json:"media_type"`
// User-Agent
UserAgent *string `json:"user_agent"`
+12 -2
View File
@@ -58,7 +58,7 @@ func TestResolvePageImagePath(t *testing.T) {
if !ok {
t.Fatal("expected direct image path to be accepted")
}
want := filepath.Join(base, "logo.png")
want := mustEvalSymlinks(t, filepath.Join(base, "logo.png"))
if got != want {
t.Fatalf("path = %q, want %q", got, want)
}
@@ -67,7 +67,7 @@ func TestResolvePageImagePath(t *testing.T) {
if !ok {
t.Fatal("expected nested image path to be accepted")
}
want = filepath.Join(base, "images", "logo.png")
want = mustEvalSymlinks(t, filepath.Join(base, "images", "logo.png"))
if got != want {
t.Fatalf("path = %q, want %q", got, want)
}
@@ -100,3 +100,13 @@ func TestResolvePageImagePathRejectsSymlinkEscape(t *testing.T) {
t.Fatalf("expected symlink escape to be rejected, got %q", got)
}
}
func mustEvalSymlinks(t *testing.T, path string) string {
t.Helper()
realPath, err := filepath.EvalSymlinks(path)
if err != nil {
t.Fatalf("eval symlinks for %q: %v", path, err)
}
return realPath
}
@@ -44,6 +44,33 @@ func TestMigrationsRunner_IsIdempotent_AndSchemaIsUpToDate(t *testing.T) {
requireColumn(t, tx, "usage_logs", "billing_type", "smallint", 0, false)
requireColumn(t, tx, "usage_logs", "request_type", "smallint", 0, false)
requireColumn(t, tx, "usage_logs", "openai_ws_mode", "boolean", 0, false)
requireColumn(t, tx, "usage_logs", "image_input_size", "character varying", 32, true)
requireColumn(t, tx, "usage_logs", "image_output_size", "character varying", 32, true)
requireColumn(t, tx, "usage_logs", "image_size_source", "character varying", 16, true)
requireColumn(t, tx, "usage_logs", "image_size_breakdown", "jsonb", 0, true)
requireConstraintDefinitionContains(
t,
tx,
"usage_logs",
"usage_logs_image_size_source_check",
"image_size_source",
"'output'",
"'input'",
"'default'",
"'legacy'",
)
requireConstraintDefinitionContains(
t,
tx,
"usage_logs",
"usage_logs_image_billing_size_check",
"image_count",
"image_size IS NOT NULL",
"'1K'",
"'2K'",
"'4K'",
"'mixed'",
)
// usage_billing_dedup: billing idempotency narrow table
var usageBillingDedupRegclass sql.NullString
+112 -13
View File
@@ -28,7 +28,7 @@ import (
gocache "github.com/patrickmn/go-cache"
)
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, 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, service_tier, reasoning_effort, inbound_endpoint, upstream_endpoint, cache_ttl_overridden, 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, 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, service_tier, reasoning_effort, inbound_endpoint, upstream_endpoint, cache_ttl_overridden, channel_id, model_mapping_chain, billing_tier, billing_mode, account_stats_cost, created_at"
// usageLogInsertArgTypes must stay in the same order as:
// 1. prepareUsageLogInsert().args
@@ -73,6 +73,10 @@ var usageLogInsertArgTypes = [...]string{
"text", // ip_address
"integer", // image_count
"text", // image_size
"text", // image_input_size
"text", // image_output_size
"text", // image_size_source
"jsonb", // image_size_breakdown
"text", // service_tier
"text", // reasoning_effort
"text", // inbound_endpoint
@@ -120,6 +124,24 @@ func appendRawUsageLogModelWhereCondition(conditions []string, args []any, model
return conditions, args
}
func appendUsageLogBillingModeWhereCondition(conditions []string, args []any, billingMode string) ([]string, []any) {
mode := strings.TrimSpace(billingMode)
if mode == "" {
return conditions, args
}
placeholder := fmt.Sprintf("$%d", len(args)+1)
switch service.BillingMode(mode) {
case service.BillingModeImage:
conditions = append(conditions, fmt.Sprintf("(billing_mode = %s OR COALESCE(image_count, 0) > 0)", placeholder))
case service.BillingModeToken:
conditions = append(conditions, fmt.Sprintf("(billing_mode = %s OR ((billing_mode IS NULL OR billing_mode = '') AND COALESCE(image_count, 0) <= 0))", placeholder))
default:
conditions = append(conditions, fmt.Sprintf("billing_mode = %s", placeholder))
}
args = append(args, mode)
return conditions, args
}
// appendRawUsageLogModelQueryFilter keeps direct model filters on the raw model column for backward
// compatibility with historical rows. Requested/upstream analytics must use
// resolveModelDimensionExpression instead.
@@ -352,6 +374,10 @@ func (r *usageLogRepository) createSingle(ctx context.Context, sqlq sqlExecutor,
ip_address,
image_count,
image_size,
image_input_size,
image_output_size,
image_size_source,
image_size_breakdown,
service_tier,
reasoning_effort,
inbound_endpoint,
@@ -369,7 +395,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
$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
)
ON CONFLICT (request_id, api_key_id) DO NOTHING
RETURNING id, created_at
@@ -790,6 +816,10 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage
ip_address,
image_count,
image_size,
image_input_size,
image_output_size,
image_size_source,
image_size_breakdown,
service_tier,
reasoning_effort,
inbound_endpoint,
@@ -803,7 +833,7 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage
created_at
) AS (VALUES `)
args := make([]any, 0, len(keys)*46)
args := make([]any, 0, len(keys)*50)
argPos := 1
for idx, key := range keys {
if idx > 0 {
@@ -867,6 +897,10 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage
ip_address,
image_count,
image_size,
image_input_size,
image_output_size,
image_size_source,
image_size_breakdown,
service_tier,
reasoning_effort,
inbound_endpoint,
@@ -915,6 +949,10 @@ func buildUsageLogBatchInsertQuery(keys []string, preparedByKey map[string]usage
ip_address,
image_count,
image_size,
image_input_size,
image_output_size,
image_size_source,
image_size_breakdown,
service_tier,
reasoning_effort,
inbound_endpoint,
@@ -1003,6 +1041,10 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) (
ip_address,
image_count,
image_size,
image_input_size,
image_output_size,
image_size_source,
image_size_breakdown,
service_tier,
reasoning_effort,
inbound_endpoint,
@@ -1016,7 +1058,7 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) (
created_at
) AS (VALUES `)
args := make([]any, 0, len(preparedList)*46)
args := make([]any, 0, len(preparedList)*50)
argPos := 1
for idx, prepared := range preparedList {
if idx > 0 {
@@ -1077,6 +1119,10 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) (
ip_address,
image_count,
image_size,
image_input_size,
image_output_size,
image_size_source,
image_size_breakdown,
service_tier,
reasoning_effort,
inbound_endpoint,
@@ -1125,6 +1171,10 @@ func buildUsageLogBestEffortInsertQuery(preparedList []usageLogInsertPrepared) (
ip_address,
image_count,
image_size,
image_input_size,
image_output_size,
image_size_source,
image_size_breakdown,
service_tier,
reasoning_effort,
inbound_endpoint,
@@ -1181,6 +1231,10 @@ func execUsageLogInsertNoResult(ctx context.Context, sqlq sqlExecutor, prepared
ip_address,
image_count,
image_size,
image_input_size,
image_output_size,
image_size_source,
image_size_breakdown,
service_tier,
reasoning_effort,
inbound_endpoint,
@@ -1198,7 +1252,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
$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
)
ON CONFLICT (request_id, api_key_id) DO NOTHING
`, prepared.args...)
@@ -1225,6 +1279,10 @@ func prepareUsageLogInsert(log *service.UsageLog) usageLogInsertPrepared {
userAgent := nullString(log.UserAgent)
ipAddress := nullString(log.IPAddress)
imageSize := nullString(log.ImageSize)
imageInputSize := nullString(log.ImageInputSize)
imageOutputSize := nullString(log.ImageOutputSize)
imageSizeSource := nullString(log.ImageSizeSource)
imageSizeBreakdown := nullStringIntMapJSON(log.ImageSizeBreakdown)
serviceTier := nullString(log.ServiceTier)
reasoningEffort := nullString(log.ReasoningEffort)
inboundEndpoint := nullString(log.InboundEndpoint)
@@ -1285,6 +1343,10 @@ func prepareUsageLogInsert(log *service.UsageLog) usageLogInsertPrepared {
ipAddress,
log.ImageCount,
imageSize,
imageInputSize,
imageOutputSize,
imageSizeSource,
imageSizeBreakdown,
serviceTier,
reasoningEffort,
inboundEndpoint,
@@ -2662,10 +2724,7 @@ func (r *usageLogRepository) ListWithFilters(ctx context.Context, params paginat
conditions = append(conditions, fmt.Sprintf("billing_type = $%d", len(args)+1))
args = append(args, int16(*filters.BillingType))
}
if filters.BillingMode != "" {
conditions = append(conditions, fmt.Sprintf("billing_mode = $%d", len(args)+1))
args = append(args, filters.BillingMode)
}
conditions, args = appendUsageLogBillingModeWhereCondition(conditions, args, filters.BillingMode)
if filters.StartTime != nil {
conditions = append(conditions, fmt.Sprintf("created_at >= $%d", len(args)+1))
args = append(args, *filters.StartTime)
@@ -3363,10 +3422,7 @@ func (r *usageLogRepository) GetStatsWithFilters(ctx context.Context, filters Us
conditions = append(conditions, fmt.Sprintf("billing_type = $%d", len(args)+1))
args = append(args, int16(*filters.BillingType))
}
if filters.BillingMode != "" {
conditions = append(conditions, fmt.Sprintf("billing_mode = $%d", len(args)+1))
args = append(args, filters.BillingMode)
}
conditions, args = appendUsageLogBillingModeWhereCondition(conditions, args, filters.BillingMode)
if filters.StartTime != nil {
conditions = append(conditions, fmt.Sprintf("created_at >= $%d", len(args)+1))
args = append(args, *filters.StartTime)
@@ -4084,6 +4140,10 @@ func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, e
ipAddress sql.NullString
imageCount int
imageSize sql.NullString
imageInputSize sql.NullString
imageOutputSize sql.NullString
imageSizeSource sql.NullString
imageSizeBreakdown sql.NullString
serviceTier sql.NullString
reasoningEffort sql.NullString
inboundEndpoint sql.NullString
@@ -4134,6 +4194,10 @@ func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, e
&ipAddress,
&imageCount,
&imageSize,
&imageInputSize,
&imageOutputSize,
&imageSizeSource,
&imageSizeBreakdown,
&serviceTier,
&reasoningEffort,
&inboundEndpoint,
@@ -4212,6 +4276,16 @@ func scanUsageLog(scanner interface{ Scan(...any) error }) (*service.UsageLog, e
if imageSize.Valid {
log.ImageSize = &imageSize.String
}
if imageInputSize.Valid {
log.ImageInputSize = &imageInputSize.String
}
if imageOutputSize.Valid {
log.ImageOutputSize = &imageOutputSize.String
}
if imageSizeSource.Valid {
log.ImageSizeSource = &imageSizeSource.String
}
log.ImageSizeBreakdown = stringIntMapFromNullJSON(imageSizeBreakdown)
if serviceTier.Valid {
log.ServiceTier = &serviceTier.String
}
@@ -4378,6 +4452,31 @@ func nullString(v *string) sql.NullString {
return sql.NullString{String: *v, Valid: true}
}
func nullStringIntMapJSON(v map[string]int) any {
if len(v) == 0 {
return nil
}
payload, err := json.Marshal(v)
if err != nil {
return nil
}
return string(payload)
}
func stringIntMapFromNullJSON(v sql.NullString) map[string]int {
if !v.Valid || strings.TrimSpace(v.String) == "" {
return nil
}
var out map[string]int
if err := json.Unmarshal([]byte(v.String), &out); err != nil {
return nil
}
if len(out) == 0 {
return nil
}
return out
}
func coalesceTrimmedString(v sql.NullString, fallback string) string {
if v.Valid && strings.TrimSpace(v.String) != "" {
return v.String
@@ -76,6 +76,10 @@ func TestUsageLogRepositoryCreateSyncRequestTypeAndLegacyFields(t *testing.T) {
sqlmock.AnyArg(), // ip_address
log.ImageCount,
sqlmock.AnyArg(), // image_size
sqlmock.AnyArg(), // image_input_size
sqlmock.AnyArg(), // image_output_size
sqlmock.AnyArg(), // image_size_source
sqlmock.AnyArg(), // image_size_breakdown
sqlmock.AnyArg(), // service_tier
sqlmock.AnyArg(), // reasoning_effort
sqlmock.AnyArg(), // inbound_endpoint
@@ -155,6 +159,10 @@ func TestUsageLogRepositoryCreate_PersistsServiceTier(t *testing.T) {
sqlmock.AnyArg(),
log.ImageCount,
sqlmock.AnyArg(),
sqlmock.AnyArg(), // image_input_size
sqlmock.AnyArg(), // image_output_size
sqlmock.AnyArg(), // image_size_source
sqlmock.AnyArg(), // image_size_breakdown
serviceTier,
sqlmock.AnyArg(),
sqlmock.AnyArg(),
@@ -230,12 +238,72 @@ func TestPrepareUsageLogInsert_ArgCountMatchesTypes(t *testing.T) {
require.Len(t, prepared.args, len(usageLogInsertArgTypes))
}
func TestPrepareUsageLogInsert_PersistsImageSizeMetadata(t *testing.T) {
imageSize := "4K"
inputSize := "1024x1024"
outputSize := "3840x2160"
source := "output"
prepared := prepareUsageLogInsert(&service.UsageLog{
UserID: 1,
APIKeyID: 2,
AccountID: 3,
RequestID: "req-image-metadata",
Model: "gpt-image-2",
RequestedModel: "gpt-image-2",
ImageCount: 2,
ImageSize: &imageSize,
ImageInputSize: &inputSize,
ImageOutputSize: &outputSize,
ImageSizeSource: &source,
ImageSizeBreakdown: map[string]int{"1K": 1, "4K": 1},
CreatedAt: time.Date(2025, 1, 6, 12, 0, 0, 0, time.UTC),
})
require.Equal(t, sql.NullString{String: imageSize, Valid: true}, prepared.args[34])
require.Equal(t, sql.NullString{String: inputSize, Valid: true}, prepared.args[35])
require.Equal(t, sql.NullString{String: outputSize, Valid: true}, prepared.args[36])
require.Equal(t, sql.NullString{String: source, Valid: true}, prepared.args[37])
require.JSONEq(t, `{"1K":1,"4K":1}`, prepared.args[38].(string))
}
func TestCoalesceTrimmedString(t *testing.T) {
require.Equal(t, "fallback", coalesceTrimmedString(sql.NullString{}, "fallback"))
require.Equal(t, "fallback", coalesceTrimmedString(sql.NullString{Valid: true, String: " "}, "fallback"))
require.Equal(t, "value", coalesceTrimmedString(sql.NullString{Valid: true, String: "value"}, "fallback"))
}
func TestAppendUsageLogBillingModeWhereCondition(t *testing.T) {
tests := []struct {
name string
billingMode string
wantCondition string
}{
{
name: "image includes legacy image rows",
billingMode: string(service.BillingModeImage),
wantCondition: "(billing_mode = $1 OR COALESCE(image_count, 0) > 0)",
},
{
name: "token includes legacy non-image rows",
billingMode: string(service.BillingModeToken),
wantCondition: "(billing_mode = $1 OR ((billing_mode IS NULL OR billing_mode = '') AND COALESCE(image_count, 0) <= 0))",
},
{
name: "per request remains exact",
billingMode: string(service.BillingModePerRequest),
wantCondition: "billing_mode = $1",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
conditions, args := appendUsageLogBillingModeWhereCondition(nil, nil, tt.billingMode)
require.Equal(t, []string{tt.wantCondition}, conditions)
require.Equal(t, []any{tt.billingMode}, args)
})
}
}
func anySliceToDriverValues(values []any) []driver.Value {
out := make([]driver.Value, 0, len(values))
for _, value := range values {
@@ -528,6 +596,63 @@ func (s usageLogScannerStub) Scan(dest ...any) error {
}
func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) {
t.Run("image_size_metadata_is_scanned", func(t *testing.T) {
now := time.Now().UTC()
log, err := scanUsageLog(usageLogScannerStub{values: []any{
int64(4),
int64(13),
int64(23),
int64(33),
sql.NullString{Valid: true, String: "req-image-metadata"},
"gpt-image-2",
sql.NullString{Valid: true, String: "gpt-image-2"},
sql.NullString{},
sql.NullInt64{},
sql.NullInt64{},
0, 0, 0, 0, 0, 0,
0, 0.0, // image_output_tokens, image_output_cost
0.0, 0.0, 0.0, 0.0, 0.8, 0.8,
1.0,
sql.NullFloat64{},
int16(service.BillingTypeBalance),
int16(service.RequestTypeSync),
false,
false,
sql.NullInt64{},
sql.NullInt64{},
sql.NullString{},
sql.NullString{},
2,
sql.NullString{Valid: true, String: "4K"},
sql.NullString{Valid: true, String: "1024x1024"},
sql.NullString{Valid: true, String: "3840x2160"},
sql.NullString{Valid: true, String: "output"},
sql.NullString{Valid: true, String: `{"4K":2}`},
sql.NullString{},
sql.NullString{},
sql.NullString{},
sql.NullString{},
false,
sql.NullInt64{},
sql.NullString{},
sql.NullString{},
sql.NullString{},
sql.NullFloat64{},
now,
}})
require.NoError(t, err)
require.Equal(t, 2, log.ImageCount)
require.NotNil(t, log.ImageSize)
require.Equal(t, "4K", *log.ImageSize)
require.NotNil(t, log.ImageInputSize)
require.Equal(t, "1024x1024", *log.ImageInputSize)
require.NotNil(t, log.ImageOutputSize)
require.Equal(t, "3840x2160", *log.ImageOutputSize)
require.NotNil(t, log.ImageSizeSource)
require.Equal(t, "output", *log.ImageSizeSource)
require.Equal(t, map[string]int{"4K": 2}, log.ImageSizeBreakdown)
})
t.Run("request_type_ws_v2_overrides_legacy", func(t *testing.T) {
now := time.Now().UTC()
log, err := scanUsageLog(usageLogScannerStub{values: []any{
@@ -567,6 +692,10 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) {
sql.NullString{},
0,
sql.NullString{},
sql.NullString{}, // image_input_size
sql.NullString{}, // image_output_size
sql.NullString{}, // image_size_source
sql.NullString{}, // image_size_breakdown
sql.NullString{Valid: true, String: "priority"},
sql.NullString{},
sql.NullString{},
@@ -615,6 +744,10 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) {
sql.NullString{},
0,
sql.NullString{},
sql.NullString{}, // image_input_size
sql.NullString{}, // image_output_size
sql.NullString{}, // image_size_source
sql.NullString{}, // image_size_breakdown
sql.NullString{Valid: true, String: "flex"},
sql.NullString{},
sql.NullString{},
@@ -663,6 +796,10 @@ func TestScanUsageLogRequestTypeAndLegacyFallback(t *testing.T) {
sql.NullString{},
0,
sql.NullString{},
sql.NullString{}, // image_input_size
sql.NullString{}, // image_output_size
sql.NullString{}, // image_size_source
sql.NullString{}, // image_size_breakdown
sql.NullString{Valid: true, String: "priority"},
sql.NullString{},
sql.NullString{},
@@ -554,6 +554,10 @@ func TestAPIContracts(t *testing.T) {
"first_token_ms": 50,
"image_count": 0,
"image_size": null,
"image_input_size": null,
"image_output_size": null,
"image_size_source": null,
"image_size_breakdown": null,
"media_type": null,
"cache_ttl_overridden": false,
"created_at": "2025-01-02T03:04:05Z",
@@ -2094,7 +2094,8 @@ func (s *AntigravityGatewayService) ForwardGemini(ctx context.Context, c *gin.Co
}
// 解析请求以获取 image_size(用于图片计费)
imageSize := s.extractImageSize(body)
imageInputSize := s.extractImageInputSize(body)
imageSize := normalizeOpenAIImageSizeTier(imageInputSize)
switch action {
case "generateContent", "streamGenerateContent":
@@ -2465,6 +2466,7 @@ handleSuccess:
ClientDisconnect: clientDisconnect,
ImageCount: imageCount,
ImageSize: imageSize,
ImageInputSize: imageInputSize,
}, nil
}
@@ -4065,19 +4067,20 @@ func (s *AntigravityGatewayService) handleClaudeStreamingResponse(c *gin.Context
// extractImageSize 从 Gemini 请求中提取 image_size 参数
func (s *AntigravityGatewayService) extractImageSize(body []byte) string {
return normalizeOpenAIImageSizeTier(s.extractImageInputSize(body))
}
func (s *AntigravityGatewayService) extractImageInputSize(body []byte) string {
var req antigravity.GeminiRequest
if err := json.Unmarshal(body, &req); err != nil {
return "2K" // 默认 2K
return ""
}
if req.GenerationConfig != nil && req.GenerationConfig.ImageConfig != nil {
size := strings.ToUpper(strings.TrimSpace(req.GenerationConfig.ImageConfig.ImageSize))
if size == "1K" || size == "2K" || size == "4K" {
return size
}
return strings.TrimSpace(req.GenerationConfig.ImageConfig.ImageSize)
}
return "2K" // 默认 2K
return ""
}
// isImageGenerationModel 判断模型是否为图片生成模型
@@ -809,6 +809,7 @@ func (s *BillingService) CalculateImageCost(model string, imageSize string, imag
if imageCount <= 0 {
return &CostBreakdown{}
}
imageSize = NormalizeImageBillingTierOrDefault(imageSize)
// 获取单价
unitPrice := s.getImageUnitPrice(model, imageSize, groupConfig)
@@ -48,6 +48,21 @@ func TestCalculateImageCost_GroupCustomPricing(t *testing.T) {
require.InDelta(t, 0.30, cost.TotalCost, 0.0001)
}
func TestCalculateImageCost_NormalizesInvalidSizeTo2K(t *testing.T) {
svc := &BillingService{}
price2K := 0.25
groupConfig := &ImagePriceConfig{Price2K: &price2K}
for _, imageSize := range []string{"", "auto", "not-a-size"} {
t.Run(imageSize, func(t *testing.T) {
cost := svc.CalculateImageCost("gemini-3-pro-image", imageSize, 2, groupConfig, 1.0)
require.InDelta(t, 0.50, cost.TotalCost, 0.0001)
require.InDelta(t, 0.50, cost.ActualCost, 0.0001)
})
}
}
// TestCalculateImageCost_4KDoublePrice 测试 4K 默认价格翻倍
func TestCalculateImageCost_4KDoublePrice(t *testing.T) {
svc := &BillingService{}
@@ -192,6 +192,46 @@ func TestGatewayServiceRecordUsage_PreservesRequestedAndUpstreamModels(t *testin
require.Equal(t, mappedModel, *usageRepo.lastLog.UpstreamModel)
}
func TestGatewayServiceRecordUsage_EmptyImageSizeDefaultsBeforeBillingAndPersistence(t *testing.T) {
imagePrice2K := 0.19
groupID := int64(901)
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
svc := newGatewayRecordUsageServiceForTest(usageRepo, &openAIRecordUsageUserRepoStub{}, &openAIRecordUsageSubRepoStub{})
err := svc.RecordUsage(context.Background(), &RecordUsageInput{
Result: &ForwardResult{
RequestID: "gateway_image_default_size",
Model: "gemini-image",
ImageCount: 1,
ImageInputSize: "auto",
Duration: time.Second,
},
APIKey: &APIKey{
ID: 801,
GroupID: i64p(groupID),
Group: &Group{
ID: groupID,
RateMultiplier: 1.0,
ImagePrice2K: &imagePrice2K,
},
},
User: &User{ID: 601},
Account: &Account{ID: 701},
})
require.NoError(t, err)
require.NotNil(t, usageRepo.lastLog)
require.Equal(t, 1, usageRepo.lastLog.ImageCount)
require.NotNil(t, usageRepo.lastLog.ImageSize)
require.Equal(t, ImageBillingSize2K, *usageRepo.lastLog.ImageSize)
require.NotNil(t, usageRepo.lastLog.ImageInputSize)
require.Equal(t, "auto", *usageRepo.lastLog.ImageInputSize)
require.NotNil(t, usageRepo.lastLog.ImageSizeSource)
require.Equal(t, ImageSizeSourceDefault, *usageRepo.lastLog.ImageSizeSource)
require.InDelta(t, 0.19, usageRepo.lastLog.TotalCost, 1e-12)
require.InDelta(t, 0.19, usageRepo.lastLog.ActualCost, 1e-12)
}
func TestGatewayServiceRecordUsage_UsageLogWriteErrorDoesNotSkipBilling(t *testing.T) {
usageRepo := &openAIRecordUsageLogRepoStub{inserted: false, err: MarkUsageLogCreateNotPersisted(context.Canceled)}
userRepo := &openAIRecordUsageUserRepoStub{}
+15 -4
View File
@@ -501,8 +501,13 @@ type ForwardResult struct {
ReasoningEffort *string
// 图片生成计费字段(图片生成模型使用)
ImageCount int // 生成的图片数量
ImageSize string // 图片尺寸 "1K", "2K", "4K"
ImageCount int // 生成的图片数量
ImageSize string // 最终计费尺寸 "1K", "2K", "4K"
ImageInputSize string // 请求中的原始图片尺寸
ImageOutputSize string // 上游响应中的图片尺寸
ImageOutputSizes []string
ImageSizeSource string
ImageSizeBreakdown map[string]int
}
// UpstreamFailoverError indicates an upstream error that should trigger account failover.
@@ -8369,6 +8374,7 @@ func (s *GatewayService) recordUsageCore(ctx context.Context, input *recordUsage
user := input.User
account := input.Account
subscription := input.Subscription
ApplyForwardImageBillingResolution(result)
// 强制缓存计费:将 input_tokens 转为 cache_read_input_tokens
// 用于粘性会话切换时的特殊计费处理
@@ -8514,6 +8520,7 @@ func (s *GatewayService) calculateImageCost(
billingModel string,
multiplier float64,
) *CostBreakdown {
sizeTier := NormalizeImageBillingTierOrDefault(result.ImageSize)
if resolved := s.resolveChannelPricing(ctx, billingModel, apiKey); resolved != nil {
tokens := UsageTokens{
InputTokens: result.Usage.InputTokens,
@@ -8527,7 +8534,7 @@ func (s *GatewayService) calculateImageCost(
GroupID: &gid,
Tokens: tokens,
RequestCount: result.ImageCount,
SizeTier: result.ImageSize,
SizeTier: sizeTier,
RateMultiplier: multiplier,
Resolver: s.resolver,
Resolved: resolved,
@@ -8547,7 +8554,7 @@ func (s *GatewayService) calculateImageCost(
Price4K: apiKey.Group.ImagePrice4K,
}
}
return s.billingService.CalculateImageCost(billingModel, result.ImageSize, result.ImageCount, groupConfig, multiplier)
return s.billingService.CalculateImageCost(billingModel, sizeTier, result.ImageCount, groupConfig, multiplier)
}
// calculateTokenCost 计算 Token 计费:根据 opts 决定走普通/长上下文/渠道统一计费。
@@ -8648,6 +8655,10 @@ func (s *GatewayService) buildRecordUsageLog(
FirstTokenMs: result.FirstTokenMs,
ImageCount: result.ImageCount,
ImageSize: optionalTrimmedStringPtr(result.ImageSize),
ImageInputSize: optionalTrimmedStringPtr(result.ImageInputSize),
ImageOutputSize: optionalTrimmedStringPtr(result.ImageOutputSize),
ImageSizeSource: optionalTrimmedStringPtr(result.ImageSizeSource),
ImageSizeBreakdown: result.ImageSizeBreakdown,
CacheTTLOverridden: cacheTTLOverridden,
ChannelID: optionalInt64Ptr(input.ChannelID),
ModelMappingChain: optionalTrimmedStringPtr(input.ModelMappingChain),
@@ -1072,21 +1072,23 @@ func (s *GeminiMessagesCompatService) Forward(ctx context.Context, c *gin.Contex
// 图片生成计费
imageCount := 0
imageSize := s.extractImageSize(body)
imageInputSize := s.extractImageInputSize(body)
imageSize := normalizeOpenAIImageSizeTier(imageInputSize)
if isImageGenerationModel(originalModel) {
imageCount = 1
}
return &ForwardResult{
RequestID: requestID,
Usage: *usage,
Model: originalModel,
UpstreamModel: mappedModel,
Stream: req.Stream,
Duration: time.Since(startTime),
FirstTokenMs: firstTokenMs,
ImageCount: imageCount,
ImageSize: imageSize,
RequestID: requestID,
Usage: *usage,
Model: originalModel,
UpstreamModel: mappedModel,
Stream: req.Stream,
Duration: time.Since(startTime),
FirstTokenMs: firstTokenMs,
ImageCount: imageCount,
ImageSize: imageSize,
ImageInputSize: imageInputSize,
}, nil
}
@@ -1600,21 +1602,23 @@ func (s *GeminiMessagesCompatService) ForwardNative(ctx context.Context, c *gin.
// 图片生成计费
imageCount := 0
imageSize := s.extractImageSize(body)
imageInputSize := s.extractImageInputSize(body)
imageSize := normalizeOpenAIImageSizeTier(imageInputSize)
if isImageGenerationModel(originalModel) {
imageCount = 1
}
return &ForwardResult{
RequestID: requestID,
Usage: *usage,
Model: originalModel,
UpstreamModel: mappedModel,
Stream: stream,
Duration: time.Since(startTime),
FirstTokenMs: firstTokenMs,
ImageCount: imageCount,
ImageSize: imageSize,
RequestID: requestID,
Usage: *usage,
Model: originalModel,
UpstreamModel: mappedModel,
Stream: stream,
Duration: time.Since(startTime),
FirstTokenMs: firstTokenMs,
ImageCount: imageCount,
ImageSize: imageSize,
ImageInputSize: imageInputSize,
}, nil
}
@@ -3432,6 +3436,10 @@ func convertClaudeGenerationConfig(req map[string]any) map[string]any {
// extractImageSize 从 Gemini 请求中提取 image_size 参数
func (s *GeminiMessagesCompatService) extractImageSize(body []byte) string {
return normalizeOpenAIImageSizeTier(s.extractImageInputSize(body))
}
func (s *GeminiMessagesCompatService) extractImageInputSize(body []byte) string {
var req struct {
GenerationConfig *struct {
ImageConfig *struct {
@@ -3440,15 +3448,12 @@ func (s *GeminiMessagesCompatService) extractImageSize(body []byte) string {
} `json:"generationConfig"`
}
if err := json.Unmarshal(body, &req); err != nil {
return "2K"
return ""
}
if req.GenerationConfig != nil && req.GenerationConfig.ImageConfig != nil {
size := strings.ToUpper(strings.TrimSpace(req.GenerationConfig.ImageConfig.ImageSize))
if size == "1K" || size == "2K" || size == "4K" {
return size
}
return strings.TrimSpace(req.GenerationConfig.ImageConfig.ImageSize)
}
return "2K"
return ""
}
@@ -0,0 +1,260 @@
package service
import (
"sort"
"strconv"
"strings"
)
const (
ImageBillingSize1K = "1K"
ImageBillingSize2K = "2K"
ImageBillingSize4K = "4K"
ImageSizeSourceOutput = "output"
ImageSizeSourceInput = "input"
ImageSizeSourceDefault = "default"
ImageSizeSourceLegacy = "legacy"
)
type ImageBillingSizeResolution struct {
BillingSize string
InputSize string
OutputSize string
Source string
Breakdown map[string]int
}
func ClassifyImageBillingTier(size string) (string, bool) {
trimmed := strings.TrimSpace(size)
normalized := strings.ToLower(trimmed)
switch normalized {
case "", "auto":
return "", false
case "1k":
return ImageBillingSize1K, true
case "2k":
return ImageBillingSize2K, true
case "4k":
return ImageBillingSize4K, true
case "2048x2048", "2048x1152":
return ImageBillingSize2K, true
case "3840x2160", "2160x3840":
return ImageBillingSize4K, true
}
width, height, ok := parseImageBillingDimensions(trimmed)
if !ok {
return "", false
}
maxEdge := width
if height > maxEdge {
maxEdge = height
}
switch {
case maxEdge <= 1024:
return ImageBillingSize1K, true
case maxEdge <= 2048:
return ImageBillingSize2K, true
default:
return ImageBillingSize4K, true
}
}
func NormalizeImageBillingTierOrDefault(size string) string {
if tier, ok := ClassifyImageBillingTier(size); ok {
return tier
}
return ImageBillingSize2K
}
func ResolveImageBillingSize(inputSize string, outputSizes []string) ImageBillingSizeResolution {
inputSize = strings.TrimSpace(inputSize)
outputSizes = compactTrimmedStrings(outputSizes)
breakdown := map[string]int{}
outputSize := firstDisplayImageOutputSize(outputSizes)
outputTier := ""
for _, output := range outputSizes {
tier, ok := ClassifyImageBillingTier(output)
if !ok {
continue
}
breakdown[tier]++
if imageTierRank(tier) > imageTierRank(outputTier) {
outputTier = tier
}
}
if outputTier != "" {
return ImageBillingSizeResolution{
BillingSize: outputTier,
InputSize: inputSize,
OutputSize: outputSize,
Source: ImageSizeSourceOutput,
Breakdown: normalizeImageSizeBreakdown(breakdown),
}
}
if tier, ok := ClassifyImageBillingTier(inputSize); ok {
return ImageBillingSizeResolution{
BillingSize: tier,
InputSize: inputSize,
OutputSize: outputSize,
Source: ImageSizeSourceInput,
}
}
return ImageBillingSizeResolution{
BillingSize: ImageBillingSize2K,
InputSize: inputSize,
OutputSize: outputSize,
Source: ImageSizeSourceDefault,
}
}
func ApplyOpenAIImageBillingResolution(result *OpenAIForwardResult) {
if result == nil || result.ImageCount <= 0 {
return
}
inputSize := strings.TrimSpace(result.ImageInputSize)
if inputSize == "" && strings.TrimSpace(result.ImageSize) != ImageBillingSize2K {
inputSize = strings.TrimSpace(result.ImageSize)
}
outputSizes := result.ImageOutputSizes
if len(outputSizes) == 0 && strings.TrimSpace(result.ImageOutputSize) != "" {
outputSizes = []string{result.ImageOutputSize}
}
resolved := ResolveImageBillingSize(inputSize, outputSizes)
applyImageBillingResolution(
&result.ImageSize,
&result.ImageInputSize,
&result.ImageOutputSize,
&result.ImageSizeSource,
&result.ImageSizeBreakdown,
resolved,
)
}
func ApplyForwardImageBillingResolution(result *ForwardResult) {
if result == nil || result.ImageCount <= 0 {
return
}
inputSize := strings.TrimSpace(result.ImageInputSize)
if inputSize == "" && strings.TrimSpace(result.ImageSize) != ImageBillingSize2K {
inputSize = strings.TrimSpace(result.ImageSize)
}
outputSizes := result.ImageOutputSizes
if len(outputSizes) == 0 && strings.TrimSpace(result.ImageOutputSize) != "" {
outputSizes = []string{result.ImageOutputSize}
}
resolved := ResolveImageBillingSize(inputSize, outputSizes)
applyImageBillingResolution(
&result.ImageSize,
&result.ImageInputSize,
&result.ImageOutputSize,
&result.ImageSizeSource,
&result.ImageSizeBreakdown,
resolved,
)
}
func applyImageBillingResolution(
billingSize *string,
inputSize *string,
outputSize *string,
source *string,
breakdown *map[string]int,
resolved ImageBillingSizeResolution,
) {
*billingSize = resolved.BillingSize
*inputSize = resolved.InputSize
*outputSize = resolved.OutputSize
*source = resolved.Source
*breakdown = resolved.Breakdown
}
func parseImageBillingDimensions(size string) (int, int, bool) {
parts := strings.Split(strings.ToLower(strings.TrimSpace(size)), "x")
if len(parts) != 2 {
return 0, 0, false
}
width, err := strconv.Atoi(strings.TrimSpace(parts[0]))
if err != nil {
return 0, 0, false
}
height, err := strconv.Atoi(strings.TrimSpace(parts[1]))
if err != nil {
return 0, 0, false
}
if width <= 0 || height <= 0 {
return 0, 0, false
}
return width, height, true
}
func compactTrimmedStrings(values []string) []string {
if len(values) == 0 {
return nil
}
out := make([]string, 0, len(values))
for _, value := range values {
trimmed := strings.TrimSpace(value)
if trimmed != "" {
out = append(out, trimmed)
}
}
return out
}
func firstDisplayImageOutputSize(outputSizes []string) string {
for _, output := range outputSizes {
if trimmed := strings.TrimSpace(output); trimmed != "" {
return trimmed
}
}
return ""
}
func imageTierRank(tier string) int {
switch strings.ToUpper(strings.TrimSpace(tier)) {
case ImageBillingSize1K:
return 1
case ImageBillingSize2K:
return 2
case ImageBillingSize4K:
return 3
default:
return 0
}
}
func normalizeImageSizeBreakdown(in map[string]int) map[string]int {
if len(in) == 0 {
return nil
}
out := make(map[string]int, len(in))
for _, tier := range []string{ImageBillingSize1K, ImageBillingSize2K, ImageBillingSize4K} {
if count := in[tier]; count > 0 {
out[tier] = count
}
}
if len(out) == 0 {
return nil
}
return out
}
func SortedImageBillingBreakdownKeys(breakdown map[string]int) []string {
keys := make([]string, 0, len(breakdown))
for key := range breakdown {
keys = append(keys, key)
}
sort.Slice(keys, func(i, j int) bool {
left, right := imageTierRank(keys[i]), imageTierRank(keys[j])
if left == right {
return keys[i] < keys[j]
}
return left < right
})
return keys
}
@@ -0,0 +1,110 @@
package service
import (
"testing"
"github.com/stretchr/testify/require"
)
func TestClassifyImageBillingTier(t *testing.T) {
tests := []struct {
name string
size string
wantTier string
wantOK bool
}{
{name: "explicit 2k square", size: "2048x2048", wantTier: "2K", wantOK: true},
{name: "explicit 2k landscape", size: "2048x1152", wantTier: "2K", wantOK: true},
{name: "explicit 4k landscape", size: "3840x2160", wantTier: "4K", wantOK: true},
{name: "explicit 4k portrait", size: "2160x3840", wantTier: "4K", wantOK: true},
{name: "long edge 1k", size: "1024X768", wantTier: "1K", wantOK: true},
{name: "long edge 2k", size: "1280x768", wantTier: "2K", wantOK: true},
{name: "long edge 4k", size: "2560x1600", wantTier: "4K", wantOK: true},
{name: "tier string 1k", size: "1k", wantTier: "1K", wantOK: true},
{name: "empty", size: "", wantOK: false},
{name: "auto", size: "auto", wantOK: false},
{name: "invalid", size: "not-a-size", wantOK: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
gotTier, gotOK := ClassifyImageBillingTier(tt.size)
require.Equal(t, tt.wantOK, gotOK)
require.Equal(t, tt.wantTier, gotTier)
})
}
}
func TestResolveImageBillingSize(t *testing.T) {
tests := []struct {
name string
inputSize string
outputSizes []string
wantBilling string
wantOutput string
wantSource string
wantBreakdown map[string]int
}{
{
name: "output wins over input",
inputSize: "1024x1024",
outputSizes: []string{"3840x2160"},
wantBilling: "4K",
wantOutput: "3840x2160",
wantSource: ImageSizeSourceOutput,
wantBreakdown: map[string]int{"4K": 1},
},
{
name: "input fallback",
inputSize: "1024x1024",
wantBilling: "1K",
wantSource: ImageSizeSourceInput,
},
{
name: "auto defaults",
inputSize: "auto",
wantBilling: "2K",
wantSource: ImageSizeSourceDefault,
},
{
name: "empty defaults",
inputSize: "",
wantBilling: "2K",
wantSource: ImageSizeSourceDefault,
},
{
name: "invalid defaults",
inputSize: "largest",
wantBilling: "2K",
wantSource: ImageSizeSourceDefault,
},
{
name: "mixed output chooses highest tier",
inputSize: "1024x1024",
outputSizes: []string{"1024x1024", "3840x2160", "1280x720"},
wantBilling: "4K",
wantOutput: "1024x1024",
wantSource: ImageSizeSourceOutput,
wantBreakdown: map[string]int{"1K": 1, "2K": 1, "4K": 1},
},
{
name: "unparseable output falls back to parseable input",
inputSize: "2048x1152",
outputSizes: []string{"auto"},
wantBilling: "2K",
wantOutput: "auto",
wantSource: ImageSizeSourceInput,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got := ResolveImageBillingSize(tt.inputSize, tt.outputSizes)
require.Equal(t, tt.wantBilling, got.BillingSize)
require.Equal(t, tt.inputSize, got.InputSize)
require.Equal(t, tt.wantOutput, got.OutputSize)
require.Equal(t, tt.wantSource, got.Source)
require.Equal(t, tt.wantBreakdown, got.Breakdown)
})
}
}
@@ -170,7 +170,21 @@ func cloneRequestMapForImageIntent(body []byte) map[string]any {
return out
}
type OpenAIResponsesImageBillingConfig struct {
Model string
SizeTier string
InputSize string
}
func resolveOpenAIResponsesImageBillingConfig(reqBody map[string]any, fallbackModel string) (string, string, error) {
cfg, err := resolveOpenAIResponsesImageBillingConfigDetailed(reqBody, fallbackModel)
if err != nil {
return "", "", err
}
return cfg.Model, cfg.SizeTier, nil
}
func resolveOpenAIResponsesImageBillingConfigDetailed(reqBody map[string]any, fallbackModel string) (OpenAIResponsesImageBillingConfig, error) {
imageModel := ""
imageSize := ""
hasImageTool := false
@@ -203,12 +217,24 @@ func resolveOpenAIResponsesImageBillingConfig(reqBody map[string]any, fallbackMo
imageModel = strings.TrimSpace(fallbackModel)
}
sizeTier := normalizeOpenAIImageSizeTier(imageSize)
return imageModel, sizeTier, nil
return OpenAIResponsesImageBillingConfig{
Model: imageModel,
SizeTier: sizeTier,
InputSize: imageSize,
}, nil
}
func resolveOpenAIResponsesImageBillingConfigFromBody(body []byte, fallbackModel string) (string, string, error) {
cfg, err := resolveOpenAIResponsesImageBillingConfigDetailedFromBody(body, fallbackModel)
if err != nil {
return "", "", err
}
return cfg.Model, cfg.SizeTier, nil
}
func resolveOpenAIResponsesImageBillingConfigDetailedFromBody(body []byte, fallbackModel string) (OpenAIResponsesImageBillingConfig, error) {
reqBody := cloneRequestMapForImageIntent(body)
return resolveOpenAIResponsesImageBillingConfig(reqBody, fallbackModel)
return resolveOpenAIResponsesImageBillingConfigDetailed(reqBody, fallbackModel)
}
func isOpenAIImageBillingModelAlias(model string) bool {
@@ -140,9 +140,10 @@ func TestResolveOpenAIResponsesImageBillingConfigDoesNotRejectUnknownSizes(t *te
func TestOpenAIImageOutputCounterDeduplicatesFinalImages(t *testing.T) {
counter := newOpenAIImageOutputCounter()
counter.AddSSEData([]byte(`{"type":"response.image_generation_call.partial_image","partial_image_b64":"abc"}`))
counter.AddSSEData([]byte(`{"type":"response.output_item.done","item":{"id":"ig_1","type":"image_generation_call","result":"final-a"}}`))
counter.AddSSEData([]byte(`{"type":"response.completed","response":{"output":[{"id":"ig_1","type":"image_generation_call","result":"final-a"},{"id":"ig_2","type":"image_generation_call","result":"final-b"}]}}`))
counter.AddSSEData([]byte(`{"type":"response.output_item.done","item":{"id":"ig_1","type":"image_generation_call","result":"final-a","size":"1024x1024"}}`))
counter.AddSSEData([]byte(`{"type":"response.completed","response":{"output":[{"id":"ig_1","type":"image_generation_call","result":"final-a"},{"id":"ig_2","type":"image_generation_call","result":"final-b","size":"3840x2160"}]}}`))
require.Equal(t, 2, counter.Count())
require.Equal(t, []string{"1024x1024", "3840x2160"}, counter.Sizes())
}
func TestOpenAIImageOutputCounterCountsImagesAPIStreamShapes(t *testing.T) {
@@ -182,3 +183,36 @@ func TestOpenAIImageOutputCounterFallsBackForInvalidMultilineSSEBody(t *testing.
)
require.Equal(t, 2, counter.Count())
}
func TestCollectOpenAIResponseImageOutputSizesFromJSONBytes(t *testing.T) {
body := []byte(`{
"output": [
{"id":"ig_1","type":"image_generation_call","result":"final-a","size":"3840x2160"},
{"id":"ig_2","type":"image_generation_call","result":"final-b","size":"1024x1024"}
]
}`)
require.Equal(t, 2, countOpenAIResponseImageOutputsFromJSONBytes(body))
require.Equal(t, []string{"3840x2160", "1024x1024"}, collectOpenAIResponseImageOutputSizesFromJSONBytes(body))
}
func TestCollectOpenAIResponseImageOutputSizesFromImagesAPIData(t *testing.T) {
body := []byte(`{
"data": [
{"b64_json":"final-a","size":"2048x1152"},
{"b64_json":"final-b","size":"2048x1152"}
]
}`)
require.Equal(t, 2, countOpenAIResponseImageOutputsFromJSONBytes(body))
require.Equal(t, []string{"2048x1152", "2048x1152"}, collectOpenAIResponseImageOutputSizesFromJSONBytes(body))
}
func TestCollectOpenAIImageOutputSizesFromSSEBody(t *testing.T) {
body := "data: {\"type\":\"response.output_item.done\",\"item\":{\"id\":\"ig_1\",\"type\":\"image_generation_call\",\"result\":\"final-a\",\"size\":\"3840x2160\"}}\n\n" +
"data: {\"type\":\"response.completed\",\"response\":{\"output\":[{\"id\":\"ig_1\",\"type\":\"image_generation_call\",\"result\":\"final-a\"},{\"id\":\"ig_2\",\"type\":\"image_generation_call\",\"result\":\"final-b\",\"size\":\"1024x1024\"}]}}\n\n" +
"data: [DONE]\n\n"
require.Equal(t, 2, countOpenAIImageOutputsFromSSEBody(body))
require.Equal(t, []string{"3840x2160", "1024x1024"}, collectOpenAIImageOutputSizesFromSSEBody(body))
}
@@ -10,12 +10,18 @@ import (
type openAIImageOutputCounter struct {
seen map[string]struct{}
seenSizes map[string]string
seenOrder []string
dataSizes []string
count int
maxDataCount int
}
func newOpenAIImageOutputCounter() *openAIImageOutputCounter {
return &openAIImageOutputCounter{seen: make(map[string]struct{})}
return &openAIImageOutputCounter{
seen: make(map[string]struct{}),
seenSizes: make(map[string]string),
}
}
func (c *openAIImageOutputCounter) Count() int {
@@ -28,6 +34,25 @@ func (c *openAIImageOutputCounter) Count() int {
return c.count
}
func (c *openAIImageOutputCounter) Sizes() []string {
if c == nil {
return nil
}
sizes := make([]string, 0, len(c.seenOrder)+len(c.dataSizes))
for _, key := range c.seenOrder {
if size := strings.TrimSpace(c.seenSizes[key]); size != "" {
sizes = append(sizes, size)
}
}
if len(sizes) == 0 && len(c.dataSizes) > 0 {
sizes = append(sizes, c.dataSizes...)
}
if len(sizes) == 0 {
return nil
}
return sizes
}
func (c *openAIImageOutputCounter) AddJSONResponse(body []byte) {
if c == nil || len(body) == 0 || !gjson.ValidBytes(body) {
return
@@ -73,10 +98,20 @@ func (c *openAIImageOutputCounter) addDataArray(data gjson.Result) {
if !data.IsArray() {
return
}
count := len(data.Array())
items := data.Array()
count := len(items)
if count > c.maxDataCount {
c.maxDataCount = count
}
sizes := make([]string, 0, len(items))
for _, item := range items {
if size := strings.TrimSpace(item.Get("size").String()); size != "" {
sizes = append(sizes, size)
}
}
if len(sizes) > 0 {
c.dataSizes = sizes
}
}
func (c *openAIImageOutputCounter) addOutputArray(output gjson.Result) {
@@ -120,10 +155,18 @@ func (c *openAIImageOutputCounter) addImageOutputItem(item gjson.Result) {
if key == "" {
return
}
size := strings.TrimSpace(item.Get("size").String())
if _, exists := c.seen[key]; exists {
if size != "" && strings.TrimSpace(c.seenSizes[key]) == "" {
c.seenSizes[key] = size
}
return
}
c.seen[key] = struct{}{}
c.seenOrder = append(c.seenOrder, key)
if size != "" {
c.seenSizes[key] = size
}
c.count++
}
@@ -142,8 +185,20 @@ func countOpenAIResponseImageOutputsFromJSONBytes(body []byte) int {
return counter.Count()
}
func collectOpenAIResponseImageOutputSizesFromJSONBytes(body []byte) []string {
counter := newOpenAIImageOutputCounter()
counter.AddJSONResponse(body)
return counter.Sizes()
}
func countOpenAIImageOutputsFromSSEBody(body string) int {
counter := newOpenAIImageOutputCounter()
counter.AddSSEBody(body)
return counter.Count()
}
func collectOpenAIImageOutputSizesFromSSEBody(body string) []string {
counter := newOpenAIImageOutputCounter()
counter.AddSSEBody(body)
return counter.Sizes()
}
@@ -1320,6 +1320,93 @@ func TestOpenAIGatewayServiceRecordUsage_ImageOnlyUsageStillPersists(t *testing.
require.Equal(t, string(BillingModeImage), *usageRepo.lastLog.BillingMode)
}
func TestOpenAIGatewayServiceRecordUsage_EmptyImageSizeDefaultsBeforeBillingAndPersistence(t *testing.T) {
imagePrice2K := 0.31
groupID := int64(1201)
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
svc := newOpenAIRecordUsageServiceForTest(usageRepo, &openAIRecordUsageUserRepoStub{}, &openAIRecordUsageSubRepoStub{}, nil)
err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{
Result: &OpenAIForwardResult{
RequestID: "resp_image_default_size",
Model: "gpt-image-2",
ImageCount: 2,
ImageSize: "",
Duration: time.Second,
},
APIKey: &APIKey{
ID: 11201,
GroupID: i64p(groupID),
Group: &Group{
ID: groupID,
RateMultiplier: 1.0,
ImagePrice2K: &imagePrice2K,
},
},
User: &User{ID: 21201},
Account: &Account{ID: 31201},
})
require.NoError(t, err)
require.NotNil(t, usageRepo.lastLog)
require.Equal(t, 2, usageRepo.lastLog.ImageCount)
require.NotNil(t, usageRepo.lastLog.ImageSize)
require.Equal(t, ImageBillingSize2K, *usageRepo.lastLog.ImageSize)
require.NotNil(t, usageRepo.lastLog.ImageSizeSource)
require.Equal(t, ImageSizeSourceDefault, *usageRepo.lastLog.ImageSizeSource)
require.Nil(t, usageRepo.lastLog.ImageInputSize)
require.Nil(t, usageRepo.lastLog.ImageOutputSize)
require.InDelta(t, 0.62, usageRepo.lastLog.TotalCost, 1e-12)
require.InDelta(t, 0.62, usageRepo.lastLog.ActualCost, 1e-12)
require.NotNil(t, usageRepo.lastLog.BillingMode)
require.Equal(t, string(BillingModeImage), *usageRepo.lastLog.BillingMode)
}
func TestOpenAIGatewayServiceRecordUsage_OutputImageSizeWinsBeforeBillingAndPersistence(t *testing.T) {
imagePrice1K := 0.11
imagePrice4K := 0.44
groupID := int64(1202)
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
svc := newOpenAIRecordUsageServiceForTest(usageRepo, &openAIRecordUsageUserRepoStub{}, &openAIRecordUsageSubRepoStub{}, nil)
err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{
Result: &OpenAIForwardResult{
RequestID: "resp_image_output_size",
Model: "gpt-image-2",
ImageCount: 1,
ImageInputSize: "1024x1024",
ImageOutputSizes: []string{"3840x2160"},
Duration: time.Second,
},
APIKey: &APIKey{
ID: 11202,
GroupID: i64p(groupID),
Group: &Group{
ID: groupID,
RateMultiplier: 1.0,
ImagePrice1K: &imagePrice1K,
ImagePrice4K: &imagePrice4K,
},
},
User: &User{ID: 21202},
Account: &Account{ID: 31202},
})
require.NoError(t, err)
require.NotNil(t, usageRepo.lastLog)
require.NotNil(t, usageRepo.lastLog.ImageSize)
require.Equal(t, ImageBillingSize4K, *usageRepo.lastLog.ImageSize)
require.NotNil(t, usageRepo.lastLog.ImageInputSize)
require.Equal(t, "1024x1024", *usageRepo.lastLog.ImageInputSize)
require.NotNil(t, usageRepo.lastLog.ImageOutputSize)
require.Equal(t, "3840x2160", *usageRepo.lastLog.ImageOutputSize)
require.NotNil(t, usageRepo.lastLog.ImageSizeSource)
require.Equal(t, ImageSizeSourceOutput, *usageRepo.lastLog.ImageSizeSource)
require.Equal(t, map[string]int{ImageBillingSize4K: 1}, usageRepo.lastLog.ImageSizeBreakdown)
require.InDelta(t, 0.44, usageRepo.lastLog.TotalCost, 1e-12)
require.InDelta(t, 0.44, usageRepo.lastLog.ActualCost, 1e-12)
}
func TestOpenAIGatewayServiceRecordUsage_ImageUsesPerImageBillingEvenWithUsageTokens(t *testing.T) {
imagePrice := 0.02
groupID := int64(12)
@@ -1641,3 +1728,42 @@ func TestGatewayServiceCalculateRecordUsageCost_ChannelImageBillingUsesSizeTier(
require.InDelta(t, 0.80, cost.TotalCost, 1e-12)
require.InDelta(t, 0.80, cost.ActualCost, 1e-12)
}
func TestGatewayServiceCalculateRecordUsageCost_ChannelImageBillingNormalizesMissingSizeTier(t *testing.T) {
groupID := int64(128)
defaultPrice := 0.10
price2K := 0.22
cache := newEmptyChannelCache()
cache.pricingByGroupModel[channelModelKey{groupID: groupID, model: "gemini-image"}] = &ChannelModelPricing{
BillingMode: BillingModeImage,
PerRequestPrice: &defaultPrice,
Intervals: []PricingInterval{{
TierLabel: "2K",
PerRequestPrice: &price2K,
}},
}
cache.channelByGroupID[groupID] = &Channel{ID: groupID, Status: StatusActive}
cache.loadedAt = time.Now()
channelService := &ChannelService{}
channelService.cache.Store(cache)
svc := &GatewayService{
billingService: NewBillingService(&config.Config{}, nil),
resolver: NewModelPricingResolver(channelService, NewBillingService(&config.Config{}, nil)),
}
cost := svc.calculateRecordUsageCost(
context.Background(),
&ForwardResult{Model: "gemini-image", ImageCount: 2, ImageSize: ""},
&APIKey{GroupID: i64p(groupID), Group: &Group{ID: groupID}},
"gemini-image",
1.0,
1.0,
nil,
)
require.NotNil(t, cost)
require.Equal(t, string(BillingModeImage), cost.BillingMode)
require.InDelta(t, 0.44, cost.TotalCost, 1e-12)
require.InDelta(t, 0.44, cost.ActualCost, 1e-12)
}
@@ -228,14 +228,19 @@ type OpenAIForwardResult struct {
ServiceTier *string
// ReasoningEffort is extracted from request body (reasoning.effort) or derived from model suffix.
// Stored for usage records display; nil means not provided / not applicable.
ReasoningEffort *string
Stream bool
OpenAIWSMode bool
ResponseHeaders http.Header
Duration time.Duration
FirstTokenMs *int
ImageCount int
ImageSize string
ReasoningEffort *string
Stream bool
OpenAIWSMode bool
ResponseHeaders http.Header
Duration time.Duration
FirstTokenMs *int
ImageCount int
ImageSize string
ImageInputSize string
ImageOutputSize string
ImageOutputSizes []string
ImageSizeSource string
ImageSizeBreakdown map[string]int
}
type OpenAIWSRetryMetricsSnapshot struct {
@@ -2416,9 +2421,10 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
}
imageBillingModel := ""
imageSizeTier := ""
imageInputSize := ""
if IsImageGenerationIntentMap(openAIResponsesEndpoint, reqModel, reqBody) {
var imageCfgErr error
imageBillingModel, imageSizeTier, imageCfgErr = resolveOpenAIResponsesImageBillingConfig(reqBody, billingModel)
imageCfg, imageCfgErr := resolveOpenAIResponsesImageBillingConfigDetailed(reqBody, billingModel)
if imageCfgErr != nil {
setOpsUpstreamError(c, http.StatusBadRequest, imageCfgErr.Error(), "")
c.JSON(http.StatusBadRequest, gin.H{
@@ -2430,6 +2436,9 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
})
return nil, imageCfgErr
}
imageBillingModel = imageCfg.Model
imageSizeTier = imageCfg.SizeTier
imageInputSize = imageCfg.InputSize
}
// Re-serialize body only if modified
@@ -2671,6 +2680,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
wsResult.UpstreamModel = upstreamModel
if wsResult.ImageCount > 0 {
wsResult.ImageSize = imageSizeTier
wsResult.ImageInputSize = imageInputSize
wsResult.BillingModel = imageBillingModel
}
return wsResult, nil
@@ -2777,6 +2787,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
var usage *OpenAIUsage
var firstTokenMs *int
imageCount := 0
var imageOutputSizes []string
if reqStream {
streamResult, err := s.handleStreamingResponse(ctx, resp, c, account, startTime, originalModel, upstreamModel)
if err != nil {
@@ -2785,6 +2796,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
usage = streamResult.usage
firstTokenMs = streamResult.firstTokenMs
imageCount = streamResult.imageCount
imageOutputSizes = streamResult.imageOutputSizes
} else {
nonStreamResult, err := s.handleNonStreamingResponse(ctx, resp, c, account, originalModel, upstreamModel)
if err != nil {
@@ -2792,6 +2804,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
}
usage = nonStreamResult.usage
imageCount = nonStreamResult.imageCount
imageOutputSizes = nonStreamResult.imageOutputSizes
}
// Extract and save Codex usage snapshot from response headers (for OAuth accounts)
@@ -2823,6 +2836,8 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco
if imageCount > 0 {
forwardResult.ImageCount = imageCount
forwardResult.ImageSize = imageSizeTier
forwardResult.ImageInputSize = imageInputSize
forwardResult.ImageOutputSizes = imageOutputSizes
forwardResult.BillingModel = imageBillingModel
}
return forwardResult, nil
@@ -2927,9 +2942,10 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
}
imageBillingModel := ""
imageSizeTier := ""
imageInputSize := ""
if IsImageGenerationIntent(openAIResponsesEndpoint, reqModel, body) {
var imageCfgErr error
imageBillingModel, imageSizeTier, imageCfgErr = resolveOpenAIResponsesImageBillingConfigFromBody(body, reqModel)
imageCfg, imageCfgErr := resolveOpenAIResponsesImageBillingConfigDetailedFromBody(body, reqModel)
if imageCfgErr != nil {
setOpsUpstreamError(c, http.StatusBadRequest, imageCfgErr.Error(), "")
c.JSON(http.StatusBadRequest, gin.H{
@@ -2941,6 +2957,9 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
})
return nil, imageCfgErr
}
imageBillingModel = imageCfg.Model
imageSizeTier = imageCfg.SizeTier
imageInputSize = imageCfg.InputSize
}
logger.LegacyPrintf("service.openai_gateway",
@@ -3026,6 +3045,7 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
var usage *OpenAIUsage
var firstTokenMs *int
imageCount := 0
var imageOutputSizes []string
if reqStream {
result, err := s.handleStreamingResponsePassthrough(ctx, resp, c, account, startTime, reqModel, upstreamPassthroughModel)
if err != nil {
@@ -3034,6 +3054,7 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
usage = result.usage
firstTokenMs = result.firstTokenMs
imageCount = result.imageCount
imageOutputSizes = result.imageOutputSizes
} else {
result, err := s.handleNonStreamingResponsePassthrough(ctx, resp, c, reqModel, upstreamPassthroughModel)
if err != nil {
@@ -3041,6 +3062,7 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
}
usage = result.usage
imageCount = result.imageCount
imageOutputSizes = result.imageOutputSizes
}
if snapshot := ParseCodexRateLimitHeaders(resp.Header); snapshot != nil {
@@ -3066,6 +3088,8 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough(
if imageCount > 0 {
forwardResult.ImageCount = imageCount
forwardResult.ImageSize = imageSizeTier
forwardResult.ImageInputSize = imageInputSize
forwardResult.ImageOutputSizes = imageOutputSizes
forwardResult.BillingModel = imageBillingModel
}
return forwardResult, nil
@@ -3361,15 +3385,17 @@ func collectOpenAIPassthroughTimeoutHeaders(h http.Header) []string {
}
type openaiStreamingResultPassthrough struct {
usage *OpenAIUsage
firstTokenMs *int
imageCount int
usage *OpenAIUsage
firstTokenMs *int
imageCount int
imageOutputSizes []string
}
type openaiNonStreamingResultPassthrough struct {
*OpenAIUsage
usage *OpenAIUsage
imageCount int
usage *OpenAIUsage
imageCount int
imageOutputSizes []string
}
func openAIStreamClientOutputStarted(c *gin.Context, localStarted bool) bool {
@@ -3539,7 +3565,12 @@ func (s *OpenAIGatewayService) handleStreamingResponsePassthrough(
needModelReplace := strings.TrimSpace(originalModel) != "" && strings.TrimSpace(mappedModel) != "" && strings.TrimSpace(originalModel) != strings.TrimSpace(mappedModel)
resultWithUsage := func() *openaiStreamingResultPassthrough {
return &openaiStreamingResultPassthrough{usage: usage, firstTokenMs: firstTokenMs, imageCount: imageCounter.Count()}
return &openaiStreamingResultPassthrough{
usage: usage,
firstTokenMs: firstTokenMs,
imageCount: imageCounter.Count(),
imageOutputSizes: imageCounter.Sizes(),
}
}
for scanner.Scan() {
@@ -3696,9 +3727,10 @@ func (s *OpenAIGatewayService) handleNonStreamingResponsePassthrough(
}
c.Data(resp.StatusCode, contentType, body)
return &openaiNonStreamingResultPassthrough{
OpenAIUsage: usage,
usage: usage,
imageCount: countOpenAIResponseImageOutputsFromJSONBytes(body),
OpenAIUsage: usage,
usage: usage,
imageCount: countOpenAIResponseImageOutputsFromJSONBytes(body),
imageOutputSizes: collectOpenAIResponseImageOutputSizesFromJSONBytes(body),
}, nil
}
@@ -3758,9 +3790,10 @@ func (s *OpenAIGatewayService) handlePassthroughSSEToJSON(resp *http.Response, c
c.Data(resp.StatusCode, contentType, body)
return &openaiNonStreamingResultPassthrough{
OpenAIUsage: usage,
usage: usage,
imageCount: countOpenAIImageOutputsFromSSEBody(bodyText),
OpenAIUsage: usage,
usage: usage,
imageCount: countOpenAIImageOutputsFromSSEBody(bodyText),
imageOutputSizes: collectOpenAIImageOutputSizesFromSSEBody(bodyText),
}, nil
}
@@ -4182,15 +4215,17 @@ func (s *OpenAIGatewayService) handleCompatErrorResponse(
// openaiStreamingResult streaming response result
type openaiStreamingResult struct {
usage *OpenAIUsage
firstTokenMs *int
imageCount int
usage *OpenAIUsage
firstTokenMs *int
imageCount int
imageOutputSizes []string
}
type openaiNonStreamingResult struct {
*OpenAIUsage
usage *OpenAIUsage
imageCount int
usage *OpenAIUsage
imageCount int
imageOutputSizes []string
}
func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp *http.Response, c *gin.Context, account *Account, startTime time.Time, originalModel, mappedModel string) (*openaiStreamingResult, error) {
@@ -4303,7 +4338,12 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp
needModelReplace := originalModel != mappedModel
resultWithUsage := func() *openaiStreamingResult {
return &openaiStreamingResult{usage: usage, firstTokenMs: firstTokenMs, imageCount: imageCounter.Count()}
return &openaiStreamingResult{
usage: usage,
firstTokenMs: firstTokenMs,
imageCount: imageCounter.Count(),
imageOutputSizes: imageCounter.Sizes(),
}
}
finalizeStream := func() (*openaiStreamingResult, error) {
if !sawTerminalEvent {
@@ -4711,9 +4751,10 @@ func (s *OpenAIGatewayService) handleNonStreamingResponse(ctx context.Context, r
c.Data(resp.StatusCode, contentType, body)
return &openaiNonStreamingResult{
OpenAIUsage: usage,
usage: usage,
imageCount: countOpenAIResponseImageOutputsFromJSONBytes(body),
OpenAIUsage: usage,
usage: usage,
imageCount: countOpenAIResponseImageOutputsFromJSONBytes(body),
imageOutputSizes: collectOpenAIResponseImageOutputSizesFromJSONBytes(body),
}, nil
}
@@ -4775,9 +4816,10 @@ func (s *OpenAIGatewayService) handleSSEToJSON(resp *http.Response, c *gin.Conte
c.Data(resp.StatusCode, contentType, body)
return &openaiNonStreamingResult{
OpenAIUsage: usage,
usage: usage,
imageCount: countOpenAIImageOutputsFromSSEBody(bodyText),
OpenAIUsage: usage,
usage: usage,
imageCount: countOpenAIImageOutputsFromSSEBody(bodyText),
imageOutputSizes: collectOpenAIImageOutputSizesFromSSEBody(bodyText),
}, nil
}
@@ -5216,6 +5258,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
user := input.User
account := input.Account
subscription := input.Subscription
ApplyOpenAIImageBillingResolution(result)
// 计算实际的新输入token(减去缓存读取的token)
// 因为 input_tokens 包含了 cache_read_tokens,而缓存读取的token不应按输入价格计费
@@ -5325,6 +5368,10 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
ImageOutputTokens: result.Usage.ImageOutputTokens,
ImageCount: result.ImageCount,
ImageSize: optionalTrimmedStringPtr(result.ImageSize),
ImageInputSize: optionalTrimmedStringPtr(result.ImageInputSize),
ImageOutputSize: optionalTrimmedStringPtr(result.ImageOutputSize),
ImageSizeSource: optionalTrimmedStringPtr(result.ImageSizeSource),
ImageSizeBreakdown: result.ImageSizeBreakdown,
}
if cost != nil {
usageLog.InputCost = cost.InputCost
@@ -5493,6 +5540,7 @@ func (s *OpenAIGatewayService) calculateOpenAIImageCost(
result *OpenAIForwardResult,
multiplier float64,
) *CostBreakdown {
sizeTier := NormalizeImageBillingTierOrDefault(result.ImageSize)
if resolved := s.resolveOpenAIChannelPricing(ctx, billingModel, apiKey); resolved != nil &&
(resolved.Mode == BillingModePerRequest || resolved.Mode == BillingModeImage) {
gid := apiKey.Group.ID
@@ -5501,7 +5549,7 @@ func (s *OpenAIGatewayService) calculateOpenAIImageCost(
Model: billingModel,
GroupID: &gid,
RequestCount: result.ImageCount,
SizeTier: result.ImageSize,
SizeTier: sizeTier,
RateMultiplier: multiplier,
Resolver: s.resolver,
Resolved: resolved,
@@ -5520,7 +5568,7 @@ func (s *OpenAIGatewayService) calculateOpenAIImageCost(
Price4K: apiKey.Group.ImagePrice4K,
}
}
return s.billingService.CalculateImageCost(billingModel, result.ImageSize, result.ImageCount, groupConfig, multiplier)
return s.billingService.CalculateImageCost(billingModel, sizeTier, result.ImageCount, groupConfig, multiplier)
}
func (s *OpenAIGatewayService) resolveOpenAIChannelPricing(ctx context.Context, billingModel string, apiKey *APIKey) *ResolvedPricing {
+55 -83
View File
@@ -532,54 +532,7 @@ func isOpenAINativeImageOption(name string) bool {
}
func normalizeOpenAIImageSizeTier(size string) string {
trimmed := strings.TrimSpace(size)
normalized := strings.ToLower(trimmed)
switch normalized {
case "", "auto":
return "2K"
case "1024x1024":
return "1K"
case "1536x1024", "1024x1536", "1792x1024", "1024x1792", "2048x2048", "2048x1152", "1152x2048":
return "2K"
case "3840x2160", "2160x3840":
return "4K"
}
width, height, ok := parseOpenAIImageSizeDimensions(trimmed)
if !ok {
return "2K"
}
return classifyUnknownOpenAIImageSizeTier(width, height)
}
const (
openAIImage2KMaxPixels = 2560 * 1440
)
func parseOpenAIImageSizeDimensions(size string) (int, int, bool) {
trimmed := strings.TrimSpace(size)
parts := strings.Split(strings.ToLower(trimmed), "x")
if len(parts) != 2 {
return 0, 0, false
}
width, err := strconv.Atoi(strings.TrimSpace(parts[0]))
if err != nil {
return 0, 0, false
}
height, err := strconv.Atoi(strings.TrimSpace(parts[1]))
if err != nil {
return 0, 0, false
}
if width <= 0 || height <= 0 {
return 0, 0, false
}
return width, height, true
}
func classifyUnknownOpenAIImageSizeTier(width int, height int) string {
if height > 0 && width > openAIImage2KMaxPixels/height {
return "4K"
}
return "2K"
return NormalizeImageBillingTierOrDefault(size)
}
func (s *OpenAIGatewayService) ForwardImages(
@@ -704,29 +657,46 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesAPIKey(
imageCount := parsed.N
var firstTokenMs *int
if parsed.Stream && isEventStreamResponse(resp.Header) {
streamUsage, streamCount, ttft, err := s.handleOpenAIImagesStreamingResponse(resp, c, startTime)
streamUsage, streamCount, streamSizes, ttft, err := s.handleOpenAIImagesStreamingResponse(resp, c, startTime)
if err != nil {
if streamCount > 0 {
return &OpenAIForwardResult{
RequestID: resp.Header.Get("x-request-id"),
Usage: streamUsage,
Model: requestModel,
UpstreamModel: upstreamModel,
Stream: parsed.Stream,
ResponseHeaders: resp.Header.Clone(),
Duration: time.Since(startTime),
FirstTokenMs: ttft,
ImageCount: streamCount,
ImageSize: parsed.SizeTier,
RequestID: resp.Header.Get("x-request-id"),
Usage: streamUsage,
Model: requestModel,
UpstreamModel: upstreamModel,
Stream: parsed.Stream,
ResponseHeaders: resp.Header.Clone(),
Duration: time.Since(startTime),
FirstTokenMs: ttft,
ImageCount: streamCount,
ImageSize: parsed.SizeTier,
ImageInputSize: parsed.Size,
ImageOutputSizes: streamSizes,
}, err
}
return nil, err
}
usage = streamUsage
imageCount = streamCount
imageOutputSizes := streamSizes
firstTokenMs = ttft
return &OpenAIForwardResult{
RequestID: resp.Header.Get("x-request-id"),
Usage: usage,
Model: requestModel,
UpstreamModel: upstreamModel,
Stream: parsed.Stream,
ResponseHeaders: resp.Header.Clone(),
Duration: time.Since(startTime),
FirstTokenMs: firstTokenMs,
ImageCount: imageCount,
ImageSize: parsed.SizeTier,
ImageInputSize: parsed.Size,
ImageOutputSizes: imageOutputSizes,
}, nil
} else {
nonStreamUsage, nonStreamCount, err := s.handleOpenAIImagesNonStreamingResponse(resp, c)
nonStreamUsage, nonStreamCount, nonStreamSizes, err := s.handleOpenAIImagesNonStreamingResponse(resp, c)
if err != nil {
return nil, err
}
@@ -734,19 +704,21 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesAPIKey(
if nonStreamCount > 0 {
imageCount = nonStreamCount
}
return &OpenAIForwardResult{
RequestID: resp.Header.Get("x-request-id"),
Usage: usage,
Model: requestModel,
UpstreamModel: upstreamModel,
Stream: parsed.Stream,
ResponseHeaders: resp.Header.Clone(),
Duration: time.Since(startTime),
FirstTokenMs: firstTokenMs,
ImageCount: imageCount,
ImageSize: parsed.SizeTier,
ImageInputSize: parsed.Size,
ImageOutputSizes: nonStreamSizes,
}, nil
}
return &OpenAIForwardResult{
RequestID: resp.Header.Get("x-request-id"),
Usage: usage,
Model: requestModel,
UpstreamModel: upstreamModel,
Stream: parsed.Stream,
ResponseHeaders: resp.Header.Clone(),
Duration: time.Since(startTime),
FirstTokenMs: firstTokenMs,
ImageCount: imageCount,
ImageSize: parsed.SizeTier,
}, nil
}
func (s *OpenAIGatewayService) buildOpenAIImagesRequest(
@@ -892,10 +864,10 @@ func cloneMultipartHeader(src textproto.MIMEHeader) textproto.MIMEHeader {
return dst
}
func (s *OpenAIGatewayService) handleOpenAIImagesNonStreamingResponse(resp *http.Response, c *gin.Context) (OpenAIUsage, int, error) {
func (s *OpenAIGatewayService) handleOpenAIImagesNonStreamingResponse(resp *http.Response, c *gin.Context) (OpenAIUsage, int, []string, error) {
body, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError)
if err != nil {
return OpenAIUsage{}, 0, err
return OpenAIUsage{}, 0, nil, err
}
responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
contentType := "application/json"
@@ -907,14 +879,14 @@ func (s *OpenAIGatewayService) handleOpenAIImagesNonStreamingResponse(resp *http
c.Data(resp.StatusCode, contentType, body)
usage, _ := extractOpenAIUsageFromJSONBytes(body)
return usage, extractOpenAIImageCountFromJSONBytes(body), nil
return usage, extractOpenAIImageCountFromJSONBytes(body), collectOpenAIResponseImageOutputSizesFromJSONBytes(body), nil
}
func (s *OpenAIGatewayService) handleOpenAIImagesStreamingResponse(
resp *http.Response,
c *gin.Context,
startTime time.Time,
) (OpenAIUsage, int, *int, error) {
) (OpenAIUsage, int, []string, *int, error) {
responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
contentType := strings.TrimSpace(resp.Header.Get("Content-Type"))
if contentType == "" {
@@ -925,7 +897,7 @@ func (s *OpenAIGatewayService) handleOpenAIImagesStreamingResponse(
flusher, ok := c.Writer.(http.Flusher)
if !ok {
return OpenAIUsage{}, 0, nil, fmt.Errorf("streaming is not supported by response writer")
return OpenAIUsage{}, 0, nil, nil, fmt.Errorf("streaming is not supported by response writer")
}
usage := OpenAIUsage{}
@@ -1010,12 +982,12 @@ func (s *OpenAIGatewayService) handleOpenAIImagesStreamingResponse(
}
if err != nil {
flushSSEEvent()
return usage, imageCounter.Count(), firstTokenMs, err
return usage, imageCounter.Count(), imageCounter.Sizes(), firstTokenMs, err
}
}
flushSSEEvent()
finalizeFallbackBody()
return usage, imageCounter.Count(), firstTokenMs, nil
return usage, imageCounter.Count(), imageCounter.Sizes(), firstTokenMs, nil
}
type readEvent struct {
@@ -1082,11 +1054,11 @@ func (s *OpenAIGatewayService) handleOpenAIImagesStreamingResponse(
if !ok {
flushSSEEvent()
finalizeFallbackBody()
return usage, imageCounter.Count(), firstTokenMs, nil
return usage, imageCounter.Count(), imageCounter.Sizes(), firstTokenMs, nil
}
if ev.err != nil {
flushSSEEvent()
return usage, imageCounter.Count(), firstTokenMs, ev.err
return usage, imageCounter.Count(), imageCounter.Sizes(), firstTokenMs, ev.err
}
processLine(ev.line)
case <-intervalCh:
@@ -1095,11 +1067,11 @@ func (s *OpenAIGatewayService) handleOpenAIImagesStreamingResponse(
continue
}
if clientDisconnected {
return usage, imageCounter.Count(), firstTokenMs, fmt.Errorf("image stream incomplete after timeout")
return usage, imageCounter.Count(), imageCounter.Sizes(), firstTokenMs, fmt.Errorf("image stream incomplete after timeout")
}
logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Images stream data interval timeout: interval=%s", streamInterval)
_ = s.writeOpenAIImagesStreamEvent(c, flusher, "error", buildOpenAIImagesStreamErrorBody(fmt.Sprintf("upstream image stream idle for %s", streamInterval)))
return usage, imageCounter.Count(), firstTokenMs, fmt.Errorf("image stream data interval timeout")
return usage, imageCounter.Count(), imageCounter.Sizes(), firstTokenMs, fmt.Errorf("image stream data interval timeout")
case <-keepaliveCh:
if clientDisconnected || time.Since(lastDownstreamWriteAt) < keepaliveInterval {
continue
@@ -72,6 +72,22 @@ func mergeOpenAIResponsesImageMeta(dst *openAIResponsesImageResult, src openAIRe
}
}
func openAIResponsesImageResultSizes(results []openAIResponsesImageResult) []string {
if len(results) == 0 {
return nil
}
sizes := make([]string, 0, len(results))
for _, result := range results {
if size := strings.TrimSpace(result.Size); size != "" {
sizes = append(sizes, size)
}
}
if len(sizes) == 0 {
return nil
}
return sizes
}
func extractOpenAIResponsesImageMetaFromLifecycleEvent(payload []byte) (openAIResponsesImageResult, int64, bool) {
switch gjson.GetBytes(payload, "type").String() {
case "response.created", "response.in_progress", "response.completed":
@@ -547,10 +563,10 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthNonStreamingResponse(
c *gin.Context,
responseFormat string,
fallbackModel string,
) (OpenAIUsage, int, error) {
) (OpenAIUsage, int, []string, error) {
body, err := ReadUpstreamResponseBody(resp.Body, s.cfg, c, openAITooLargeError)
if err != nil {
return OpenAIUsage{}, 0, err
return OpenAIUsage{}, 0, nil, err
}
var usage OpenAIUsage
@@ -559,10 +575,10 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthNonStreamingResponse(
})
results, createdAt, usageRaw, firstMeta, _, err := collectOpenAIImagesFromResponsesBody(body)
if err != nil {
return OpenAIUsage{}, 0, err
return OpenAIUsage{}, 0, nil, err
}
if len(results) == 0 {
return OpenAIUsage{}, 0, fmt.Errorf("upstream did not return image output")
return OpenAIUsage{}, 0, nil, fmt.Errorf("upstream did not return image output")
}
if strings.TrimSpace(firstMeta.Model) == "" {
firstMeta.Model = strings.TrimSpace(fallbackModel)
@@ -570,11 +586,11 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthNonStreamingResponse(
responseBody, err := buildOpenAIImagesAPIResponse(results, createdAt, usageRaw, firstMeta, responseFormat)
if err != nil {
return OpenAIUsage{}, 0, err
return OpenAIUsage{}, 0, nil, err
}
responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
c.Data(resp.StatusCode, "application/json; charset=utf-8", responseBody)
return usage, len(results), nil
return usage, len(results), openAIResponsesImageResultSizes(results), nil
}
func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse(
@@ -584,7 +600,7 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse(
responseFormat string,
streamPrefix string,
fallbackModel string,
) (OpenAIUsage, int, *int, error) {
) (OpenAIUsage, int, []string, *int, error) {
responseheaders.WriteFilteredHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
c.Header("Content-Type", "text/event-stream")
c.Header("Cache-Control", "no-cache")
@@ -593,7 +609,7 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse(
flusher, ok := c.Writer.(http.Flusher)
if !ok {
return OpenAIUsage{}, 0, nil, fmt.Errorf("streaming is not supported by response writer")
return OpenAIUsage{}, 0, nil, nil, fmt.Errorf("streaming is not supported by response writer")
}
format := strings.ToLower(strings.TrimSpace(responseFormat))
@@ -603,6 +619,7 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse(
usage := OpenAIUsage{}
imageCount := 0
var imageOutputSizes []string
var firstTokenMs *int
emitted := make(map[string]struct{})
pendingResults := make([]openAIResponsesImageResult, 0, 1)
@@ -713,6 +730,7 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse(
s.tryWriteOpenAIImagesStreamEvent(c, flusher, &clientDisconnected, &lastDownstreamWriteAt, eventName, payload)
}
imageCount = len(emitted)
imageOutputSizes = openAIResponsesImageResultSizes(finalResults)
processDataDone = true
}
}
@@ -753,6 +771,7 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse(
s.tryWriteOpenAIImagesStreamEvent(c, flusher, &clientDisconnected, &lastDownstreamWriteAt, eventName, payload)
}
imageCount = len(emitted)
imageOutputSizes = openAIResponsesImageResultSizes(pendingResults)
return nil
}
@@ -769,33 +788,33 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse(
line, err := reader.ReadBytes('\n')
done, processErr := processLine(line)
if processErr != nil {
return usage, imageCount, firstTokenMs, processErr
return usage, imageCount, imageOutputSizes, firstTokenMs, processErr
}
if done {
return usage, imageCount, firstTokenMs, nil
return usage, imageCount, imageOutputSizes, firstTokenMs, nil
}
if err == io.EOF {
break
}
if err != nil {
if done, processErr := flushData(); processErr != nil {
return usage, imageCount, firstTokenMs, processErr
return usage, imageCount, imageOutputSizes, firstTokenMs, processErr
} else if done {
return usage, imageCount, firstTokenMs, nil
return usage, imageCount, imageOutputSizes, firstTokenMs, nil
}
s.tryWriteOpenAIImagesStreamEvent(c, flusher, &clientDisconnected, &lastDownstreamWriteAt, "error", buildOpenAIImagesStreamErrorBody(err.Error()))
return usage, imageCount, firstTokenMs, err
return usage, imageCount, imageOutputSizes, firstTokenMs, err
}
}
if done, processErr := flushData(); processErr != nil {
return usage, imageCount, firstTokenMs, processErr
return usage, imageCount, imageOutputSizes, firstTokenMs, processErr
} else if done {
return usage, imageCount, firstTokenMs, nil
return usage, imageCount, imageOutputSizes, firstTokenMs, nil
}
if err := finalizePending(); err != nil {
return usage, imageCount, firstTokenMs, err
return usage, imageCount, imageOutputSizes, firstTokenMs, err
}
return usage, imageCount, firstTokenMs, nil
return usage, imageCount, imageOutputSizes, firstTokenMs, nil
}
type readEvent struct {
@@ -861,30 +880,30 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse(
case ev, ok := <-events:
if !ok {
if done, processErr := flushData(); processErr != nil {
return usage, imageCount, firstTokenMs, processErr
return usage, imageCount, imageOutputSizes, firstTokenMs, processErr
} else if done {
return usage, imageCount, firstTokenMs, nil
return usage, imageCount, imageOutputSizes, firstTokenMs, nil
}
if err := finalizePending(); err != nil {
return usage, imageCount, firstTokenMs, err
return usage, imageCount, imageOutputSizes, firstTokenMs, err
}
return usage, imageCount, firstTokenMs, nil
return usage, imageCount, imageOutputSizes, firstTokenMs, nil
}
if ev.err != nil {
if done, processErr := flushData(); processErr != nil {
return usage, imageCount, firstTokenMs, processErr
return usage, imageCount, imageOutputSizes, firstTokenMs, processErr
} else if done {
return usage, imageCount, firstTokenMs, nil
return usage, imageCount, imageOutputSizes, firstTokenMs, nil
}
s.tryWriteOpenAIImagesStreamEvent(c, flusher, &clientDisconnected, &lastDownstreamWriteAt, "error", buildOpenAIImagesStreamErrorBody(ev.err.Error()))
return usage, imageCount, firstTokenMs, ev.err
return usage, imageCount, imageOutputSizes, firstTokenMs, ev.err
}
done, processErr := processLine(ev.line)
if processErr != nil {
return usage, imageCount, firstTokenMs, processErr
return usage, imageCount, imageOutputSizes, firstTokenMs, processErr
}
if done {
return usage, imageCount, firstTokenMs, nil
return usage, imageCount, imageOutputSizes, firstTokenMs, nil
}
case <-intervalCh:
lastRead := time.Unix(0, atomic.LoadInt64(&lastReadAt))
@@ -892,11 +911,11 @@ func (s *OpenAIGatewayService) handleOpenAIImagesOAuthStreamingResponse(
continue
}
if clientDisconnected {
return usage, imageCount, firstTokenMs, fmt.Errorf("image stream incomplete after timeout")
return usage, imageCount, imageOutputSizes, firstTokenMs, fmt.Errorf("image stream incomplete after timeout")
}
logger.LegacyPrintf("service.openai_gateway", "[OpenAI] Images responses stream data interval timeout: interval=%s", streamInterval)
s.tryWriteOpenAIImagesStreamEvent(c, flusher, &clientDisconnected, &lastDownstreamWriteAt, "error", buildOpenAIImagesStreamErrorBody(fmt.Sprintf("upstream image stream idle for %s", streamInterval)))
return usage, imageCount, firstTokenMs, fmt.Errorf("image stream data interval timeout")
return usage, imageCount, imageOutputSizes, firstTokenMs, fmt.Errorf("image stream data interval timeout")
case <-keepaliveCh:
if clientDisconnected || time.Since(lastDownstreamWriteAt) < keepaliveInterval {
continue
@@ -1019,31 +1038,34 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesOAuth(
defer func() { _ = resp.Body.Close() }()
var (
usage OpenAIUsage
imageCount int
firstTokenMs *int
usage OpenAIUsage
imageCount int
imageOutputSizes []string
firstTokenMs *int
)
if parsed.Stream {
usage, imageCount, firstTokenMs, err = s.handleOpenAIImagesOAuthStreamingResponse(resp, c, startTime, parsed.ResponseFormat, openAIImagesStreamPrefix(parsed), requestModel)
usage, imageCount, imageOutputSizes, firstTokenMs, err = s.handleOpenAIImagesOAuthStreamingResponse(resp, c, startTime, parsed.ResponseFormat, openAIImagesStreamPrefix(parsed), requestModel)
if err != nil {
if imageCount > 0 {
return &OpenAIForwardResult{
RequestID: resp.Header.Get("x-request-id"),
Usage: usage,
Model: requestModel,
UpstreamModel: requestModel,
Stream: parsed.Stream,
ResponseHeaders: resp.Header.Clone(),
Duration: time.Since(startTime),
FirstTokenMs: firstTokenMs,
ImageCount: imageCount,
ImageSize: parsed.SizeTier,
RequestID: resp.Header.Get("x-request-id"),
Usage: usage,
Model: requestModel,
UpstreamModel: requestModel,
Stream: parsed.Stream,
ResponseHeaders: resp.Header.Clone(),
Duration: time.Since(startTime),
FirstTokenMs: firstTokenMs,
ImageCount: imageCount,
ImageSize: parsed.SizeTier,
ImageInputSize: parsed.Size,
ImageOutputSizes: imageOutputSizes,
}, err
}
return nil, err
}
} else {
usage, imageCount, err = s.handleOpenAIImagesOAuthNonStreamingResponse(resp, c, parsed.ResponseFormat, requestModel)
usage, imageCount, imageOutputSizes, err = s.handleOpenAIImagesOAuthNonStreamingResponse(resp, c, parsed.ResponseFormat, requestModel)
if err != nil {
return nil, err
}
@@ -1052,15 +1074,17 @@ func (s *OpenAIGatewayService) forwardOpenAIImagesOAuth(
imageCount = parsed.N
}
return &OpenAIForwardResult{
RequestID: resp.Header.Get("x-request-id"),
Usage: usage,
Model: requestModel,
UpstreamModel: requestModel,
Stream: parsed.Stream,
ResponseHeaders: resp.Header.Clone(),
Duration: time.Since(startTime),
FirstTokenMs: firstTokenMs,
ImageCount: imageCount,
ImageSize: parsed.SizeTier,
RequestID: resp.Header.Get("x-request-id"),
Usage: usage,
Model: requestModel,
UpstreamModel: requestModel,
Stream: parsed.Stream,
ResponseHeaders: resp.Header.Clone(),
Duration: time.Since(startTime),
FirstTokenMs: firstTokenMs,
ImageCount: imageCount,
ImageSize: parsed.SizeTier,
ImageInputSize: parsed.Size,
ImageOutputSizes: imageOutputSizes,
}, nil
}
@@ -149,9 +149,9 @@ func TestOpenAIGatewayServiceParseOpenAIImagesRequest_NormalizesOfficialAndCusto
{size: "2048x1152", wantTier: "2K"},
{size: "3840x2160", wantTier: "4K"},
{size: "2160x3840", wantTier: "4K"},
{size: "1024X768", wantTier: "2K"},
{size: "1024X768", wantTier: "1K"},
{size: "1280x768", wantTier: "2K"},
{size: "2560x1440", wantTier: "2K"},
{size: "2560x1440", wantTier: "4K"},
{size: "2560x1600", wantTier: "4K"},
{size: "auto", wantTier: "2K"},
}
@@ -186,7 +186,7 @@ func TestOpenAIGatewayServiceParseOpenAIImagesRequest_UnknownSizesDoNotBlockPass
{size: "2048x1153", wantTier: "2K"},
{size: "4096x1024", wantTier: "4K"},
{size: "3840x1024", wantTier: "4K"},
{size: "512x512", wantTier: "2K"},
{size: "512x512", wantTier: "1K"},
{size: "invalid", wantTier: "2K"},
{size: "999999999999999999999999999x2", wantTier: "2K"},
}
+26 -15
View File
@@ -2351,18 +2351,19 @@ func (s *OpenAIGatewayService) forwardOpenAIWSV2(
)
return &OpenAIForwardResult{
RequestID: responseID,
Usage: *usage,
Model: originalModel,
UpstreamModel: mappedModel,
ImageCount: imageCounter.Count(),
ServiceTier: extractOpenAIServiceTier(reqBody),
ReasoningEffort: extractOpenAIReasoningEffort(reqBody, originalModel),
Stream: reqStream,
OpenAIWSMode: true,
ResponseHeaders: lease.HandshakeHeaders(),
Duration: time.Since(startTime),
FirstTokenMs: firstTokenMs,
RequestID: responseID,
Usage: *usage,
Model: originalModel,
UpstreamModel: mappedModel,
ImageCount: imageCounter.Count(),
ImageOutputSizes: imageCounter.Sizes(),
ServiceTier: extractOpenAIServiceTier(reqBody),
ReasoningEffort: extractOpenAIReasoningEffort(reqBody, originalModel),
Stream: reqStream,
OpenAIWSMode: true,
ResponseHeaders: lease.HandshakeHeaders(),
Duration: time.Since(startTime),
FirstTokenMs: firstTokenMs,
}, nil
}
@@ -2464,6 +2465,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
originalModel string
imageBillingModel string
imageSizeTier string
imageInputSize string
payloadBytes int
}
@@ -2567,12 +2569,16 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
}
imageBillingModel := ""
imageSizeTier := ""
imageInputSize := ""
if imageIntent {
var imageCfgErr error
imageBillingModel, imageSizeTier, imageCfgErr = resolveOpenAIResponsesImageBillingConfigFromBody(normalized, originalModel)
imageCfg, imageCfgErr := resolveOpenAIResponsesImageBillingConfigDetailedFromBody(normalized, originalModel)
if imageCfgErr != nil {
return openAIWSClientPayload{}, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, imageCfgErr.Error(), imageCfgErr)
}
imageBillingModel = imageCfg.Model
imageSizeTier = imageCfg.SizeTier
imageInputSize = imageCfg.InputSize
}
// Apply OpenAI Fast Policy on the response.create frame using the same
@@ -2621,6 +2627,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
originalModel: originalModel,
imageBillingModel: imageBillingModel,
imageSizeTier: imageSizeTier,
imageInputSize: imageInputSize,
payloadBytes: len(normalized),
}, nil
}
@@ -2822,7 +2829,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
return payload, nil
}
sendAndRelay := func(turn int, lease *openAIWSConnLease, payload []byte, payloadBytes int, originalModel string, imageBillingModel string, imageSizeTier string) (*OpenAIForwardResult, error) {
sendAndRelay := func(turn int, lease *openAIWSConnLease, payload []byte, payloadBytes int, originalModel string, imageBillingModel string, imageSizeTier string, imageInputSize string) (*OpenAIForwardResult, error) {
if lease == nil {
return nil, errors.New("upstream websocket lease is nil")
}
@@ -3046,6 +3053,8 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
if imageCount > 0 {
result.ImageCount = imageCount
result.ImageSize = imageSizeTier
result.ImageInputSize = imageInputSize
result.ImageOutputSizes = imageCounter.Sizes()
result.BillingModel = imageBillingModel
}
return result, nil
@@ -3057,6 +3066,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
currentOriginalModel := firstPayload.originalModel
currentImageBillingModel := firstPayload.imageBillingModel
currentImageSizeTier := firstPayload.imageSizeTier
currentImageInputSize := firstPayload.imageInputSize
currentPayloadBytes := firstPayload.payloadBytes
isStrictAffinityTurn := func(payload []byte) bool {
if !storeDisabled {
@@ -3534,7 +3544,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
)
}
result, relayErr := sendAndRelay(turn, sessionLease, currentPayload, currentPayloadBytes, currentOriginalModel, currentImageBillingModel, currentImageSizeTier)
result, relayErr := sendAndRelay(turn, sessionLease, currentPayload, currentPayloadBytes, currentOriginalModel, currentImageBillingModel, currentImageSizeTier, currentImageInputSize)
if relayErr != nil {
lastTurnClean = false
if recoverIngressPrevResponseNotFound(relayErr, turn, connID) {
@@ -3658,6 +3668,7 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient(
currentOriginalModel = nextPayload.originalModel
currentImageBillingModel = nextPayload.imageBillingModel
currentImageSizeTier = nextPayload.imageSizeTier
currentImageInputSize = nextPayload.imageInputSize
currentPayloadBytes = nextPayload.payloadBytes
storeDisabled = s.isOpenAIWSStoreDisabledInRequestRaw(currentPayload, account)
if !storeDisabled {
+7 -3
View File
@@ -162,9 +162,13 @@ type UsageLog struct {
CacheTTLOverridden bool
// 图片生成字段
ImageCount int
ImageSize *string
MediaType *string
ImageCount int
ImageSize *string
ImageInputSize *string
ImageOutputSize *string
ImageSizeSource *string
ImageSizeBreakdown map[string]int
MediaType *string
CreatedAt time.Time