Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 17 additions & 3 deletions relay/common/relay_info.go
Original file line number Diff line number Diff line change
Expand Up @@ -46,9 +46,18 @@ type RerankerInfo struct {
}

type BuildInToolInfo struct {
ToolName string
CallCount int
SearchContextSize string
ToolName string
CallCount int
SearchContextSize string
ImageGenerationQuality string
ImageGenerationSize string
ImageGenerationTiers map[ImageGenerationTier]int
}

// ImageGenerationTier identifies one billable image quality and size combination.
type ImageGenerationTier struct {
Quality string
Size string
}

type ResponsesUsageInfo struct {
Expand Down Expand Up @@ -414,6 +423,11 @@ func GenRelayInfoResponses(c *gin.Context, request *dto.OpenAIResponsesRequest)
searchContextSize = "medium"
}
info.ResponsesUsageInfo.BuiltInTools[toolType].SearchContextSize = searchContextSize
case dto.BuildInToolImageGeneration:
info.ResponsesUsageInfo.BuiltInTools[toolType].ImageGenerationQuality =
common.Interface2String(tool["quality"])
info.ResponsesUsageInfo.BuiltInTools[toolType].ImageGenerationSize =
common.Interface2String(tool["size"])
}
}
}
Expand Down
21 changes: 21 additions & 0 deletions relay/common/relay_info_test.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package common

import (
"context"
"encoding/json"
"net/http/httptest"
"testing"
Expand Down Expand Up @@ -45,6 +46,26 @@ func TestRelayInfoGetFinalRequestRelayFormatNilReceiver(t *testing.T) {
require.Equal(t, types.RelayFormat(""), info.GetFinalRequestRelayFormat())
}

func TestGenRelayInfoResponsesCapturesImageGenerationOptions(t *testing.T) {
gin.SetMode(gin.TestMode)
ctx, _ := gin.CreateTestContext(httptest.NewRecorder())
ctx.Request = httptest.NewRequestWithContext(context.Background(), "POST", "/v1/responses", nil)

request := &dto.OpenAIResponsesRequest{
Model: "gpt-5.5",
Tools: json.RawMessage(`[
{"type":"image_generation","quality":"medium","size":"1024x1536"}
]`),
}

info := GenRelayInfoResponses(ctx, request)
tool := info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration]

require.NotNil(t, tool)
assert.Equal(t, "medium", tool.ImageGenerationQuality)
assert.Equal(t, "1024x1536", tool.ImageGenerationSize)
}

func TestRelayInfoMetaTypedNilReceiver(t *testing.T) {
var info *RelayInfo
var meta convmeta.Meta = info
Expand Down
68 changes: 61 additions & 7 deletions relay/common/tool_usage.go
Original file line number Diff line number Diff line change
Expand Up @@ -77,8 +77,9 @@ func (info *RelayInfo) incrementBillableToolCall(name string) {
// ImageGenerationCallCounter counts completed Responses image_generation_call
// outputs with stream-safe identity deduplication.
type ImageGenerationCallCounter struct {
seen map[string]struct{}
seen map[string]int
count int
tiers []ImageGenerationTier
}

// Observe records one completed final image output when billable.
Expand All @@ -97,6 +98,14 @@ func (c *ImageGenerationCallCounter) Observe(item *dto.ResponsesOutput, outputIn
case "failed", "cancelled", "canceled", "incomplete", "partial":
return
}
quality := strings.ToLower(strings.TrimSpace(item.Quality))
if quality == "auto" {
quality = ""
}
size := strings.ToLower(strings.TrimSpace(item.Size))
if size == "auto" {
size = ""
}

aliases := make([]string, 0, 4)
if item.ID != "" {
Expand All @@ -112,16 +121,33 @@ func (c *ImageGenerationCallCounter) Observe(item *dto.ResponsesOutput, outputIn
aliases = append(aliases, "result:"+hex.EncodeToString(sum[:]))

if c.seen == nil {
c.seen = make(map[string]struct{})
c.seen = make(map[string]int)
}
observedIndex := -1
for _, alias := range aliases {
if _, ok := c.seen[alias]; ok {
return
if index, ok := c.seen[alias]; ok {
observedIndex = index
break
}
}
if observedIndex >= 0 {
for _, alias := range aliases {
c.seen[alias] = observedIndex
}
if quality != "" {
c.tiers[observedIndex].Quality = quality
}
if size != "" {
c.tiers[observedIndex].Size = size
}
return
}

observedIndex = len(c.tiers)
for _, alias := range aliases {
c.seen[alias] = struct{}{}
c.seen[alias] = observedIndex
}
c.tiers = append(c.tiers, ImageGenerationTier{Quality: quality, Size: size})
c.count++
}

Expand All @@ -132,6 +158,7 @@ func (c *ImageGenerationCallCounter) Reset() {
}
c.seen = nil
c.count = 0
c.tiers = nil
}

// Count returns the deduplicated completed image output count before commit capping.
Expand Down Expand Up @@ -165,13 +192,40 @@ func (c *ImageGenerationCallCounter) Commit(info *RelayInfo) {
count = dto.MaxImageN
}

requestQuality, requestSize := "", ""
if existing, ok := info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration]; ok && existing != nil {
requestQuality = strings.ToLower(strings.TrimSpace(existing.ImageGenerationQuality))
requestSize = strings.ToLower(strings.TrimSpace(existing.ImageGenerationSize))
}
if requestQuality == "auto" {
requestQuality = ""
}
if requestSize == "auto" {
requestSize = ""
}

tierCounts := make(map[ImageGenerationTier]int)
if c != nil {
for _, tier := range c.tiers[:count] {
if tier.Quality == "" {
tier.Quality = requestQuality
}
if tier.Size == "" {
tier.Size = requestSize
}
tierCounts[tier]++
}
}

if existing, ok := info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration]; ok && existing != nil {
existing.CallCount = count
existing.ImageGenerationTiers = tierCounts
return
}
info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration] = &BuildInToolInfo{
ToolName: dto.BuildInToolImageGeneration,
CallCount: count,
ToolName: dto.BuildInToolImageGeneration,
CallCount: count,
ImageGenerationTiers: tierCounts,
}
Comment thread
coderabbitai[bot] marked this conversation as resolved.
}

Expand Down
118 changes: 117 additions & 1 deletion relay/common/tool_usage_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -316,7 +316,13 @@ func TestImageGenerationCallCounterCommitCapsAtMaxImageN(t *testing.T) {
info := &RelayInfo{}
counter.Commit(info)
require.Contains(t, info.ResponsesUsageInfo.BuiltInTools, dto.BuildInToolImageGeneration)
assert.Equal(t, dto.MaxImageN, info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration].CallCount)
tool := info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration]
assert.Equal(t, dto.MaxImageN, tool.CallCount)
tierTotal := 0
for _, count := range tool.ImageGenerationTiers {
tierTotal += count
}
assert.Equal(t, dto.MaxImageN, tierTotal)
}

func TestImageGenerationCallCounterCommitDoesNotBillDeclarationsAlone(t *testing.T) {
Expand All @@ -336,6 +342,116 @@ func TestImageGenerationCallCounterCommitDoesNotBillDeclarationsAlone(t *testing
assert.Equal(t, 0, info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration].CallCount)
}

func TestImageGenerationCallCounterCommitPrefersResponseOptions(t *testing.T) {
t.Parallel()

info := &RelayInfo{
ResponsesUsageInfo: &ResponsesUsageInfo{
BuiltInTools: map[string]*BuildInToolInfo{
dto.BuildInToolImageGeneration: {
ToolName: dto.BuildInToolImageGeneration,
ImageGenerationQuality: "medium",
ImageGenerationSize: "1024x1536",
},
},
},
}
counter := &ImageGenerationCallCounter{}
counter.Observe(&dto.ResponsesOutput{
Type: dto.ResponsesOutputTypeImageGenerationCall,
ID: "img_1",
Status: "completed",
Result: "base64-a",
}, nil)
counter.Observe(&dto.ResponsesOutput{
Type: dto.ResponsesOutputTypeImageGenerationCall,
ID: "img_1",
Status: "completed",
Result: "base64-a",
Quality: "high",
Size: "1536x1024",
}, nil)
counter.Commit(info)

tool := info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration]
require.NotNil(t, tool)
assert.Equal(t, 1, tool.CallCount)
assert.Equal(t, 1, tool.ImageGenerationTiers[ImageGenerationTier{
Quality: "high",
Size: "1536x1024",
}])
}

func TestImageGenerationCallCounterCommitKeepsRequestOptionsWhenResponseOmitsThem(t *testing.T) {
t.Parallel()

info := &RelayInfo{
ResponsesUsageInfo: &ResponsesUsageInfo{
BuiltInTools: map[string]*BuildInToolInfo{
dto.BuildInToolImageGeneration: {
ToolName: dto.BuildInToolImageGeneration,
ImageGenerationQuality: "medium",
ImageGenerationSize: "1024x1536",
},
},
},
}
counter := &ImageGenerationCallCounter{}
counter.Observe(&dto.ResponsesOutput{
Type: dto.ResponsesOutputTypeImageGenerationCall,
Status: "completed",
Result: "base64-a",
Quality: "auto",
Size: "",
}, nil)
counter.Commit(info)

tool := info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration]
require.NotNil(t, tool)
assert.Equal(t, 1, tool.CallCount)
assert.Equal(t, 1, tool.ImageGenerationTiers[ImageGenerationTier{
Quality: "medium",
Size: "1024x1536",
}])
}

func TestImageGenerationCallCounterCommitKeepsDistinctResponseTiers(t *testing.T) {
t.Parallel()

counter := &ImageGenerationCallCounter{}
counter.Observe(&dto.ResponsesOutput{
Type: dto.ResponsesOutputTypeImageGenerationCall,
ID: "img_1",
Status: "completed",
Result: "base64-a",
Quality: "high",
Size: "1024x1024",
}, nil)
counter.Observe(&dto.ResponsesOutput{
Type: dto.ResponsesOutputTypeImageGenerationCall,
ID: "img_2",
Status: "completed",
Result: "base64-b",
Quality: "high",
Size: "1024x1536",
}, nil)

info := &RelayInfo{}
counter.Commit(info)
tool := info.ResponsesUsageInfo.BuiltInTools[dto.BuildInToolImageGeneration]

require.NotNil(t, tool)
assert.Equal(t, 2, tool.CallCount)
assert.Equal(t, 1, tool.ImageGenerationTiers[ImageGenerationTier{
Quality: "high",
Size: "1024x1024",
}])
assert.Equal(t, 1, tool.ImageGenerationTiers[ImageGenerationTier{
Quality: "high",
Size: "1024x1536",
}])
}

func TestIsNonBillableResponsesStatus(t *testing.T) {
t.Parallel()

Expand Down
28 changes: 27 additions & 1 deletion service/text_quota.go
Original file line number Diff line number Diff line change
Expand Up @@ -100,11 +100,17 @@ func isLegacyClaudeDerivedOpenAIUsage(relayInfo *relaycommon.RelayInfo, usage *d
return usage.ClaudeCacheCreation5mTokens > 0 || usage.ClaudeCacheCreation1hTokens > 0
}

// collectToolSurchargeItem resolves a model-aware tool price before appending it.
func collectToolSurchargeItem(items []ToolSurchargeItem, name string, count int, modelName string) []ToolSurchargeItem {
price := operation_setting.GetToolPriceForModel(name, modelName)
return collectToolSurchargeItemWithPrice(items, name, count, price)
}

// collectToolSurchargeItemWithPrice appends one validated, explicitly priced tool charge.
func collectToolSurchargeItemWithPrice(items []ToolSurchargeItem, name string, count int, price float64) []ToolSurchargeItem {
if count <= 0 {
return items
}
price := operation_setting.GetToolPriceForModel(name, modelName)
if price <= 0 || math.IsNaN(price) || math.IsInf(price, 0) {
return items
}
Expand Down Expand Up @@ -156,6 +162,26 @@ func calculateTextToolCallSurcharge(ctx *gin.Context, relayInfo *relaycommon.Rel
if tool == nil {
continue
}
if name == dto.BuildInToolImageGeneration {
if len(tool.ImageGenerationTiers) == 0 {
price := operation_setting.GetImageGenerationToolPriceForModel(
summary.ModelName,
tool.ImageGenerationQuality,
tool.ImageGenerationSize,
)
items = collectToolSurchargeItemWithPrice(items, name, tool.CallCount, price)
continue
}
for tier, count := range tool.ImageGenerationTiers {
price := operation_setting.GetImageGenerationToolPriceForModel(
summary.ModelName,
tier.Quality,
tier.Size,
)
items = collectToolSurchargeItemWithPrice(items, name, count, price)
}
continue
}
items = collectToolSurchargeItem(items, name, tool.CallCount, summary.ModelName)
}
}
Expand Down
Loading