mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
Fix image billing size normalization
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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"`
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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{}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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"},
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user