Skip to content

Commit f2af8a0

Browse files
tools cost
1 parent e10d38b commit f2af8a0

4 files changed

Lines changed: 98 additions & 26 deletions

File tree

internal/provider/openai/cost.go

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -189,6 +189,17 @@ var OpenAiPerThousandTokenCost = map[string]map[string]float64{
189189
},
190190
}
191191

192+
var OpenAiPerThousandCallsToolCost = map[string]float64{
193+
"web_search": 10.0,
194+
"web_search_preview": 25.0,
195+
"web_search_preview_reasoning": 10.0,
196+
}
197+
198+
var AllowedTools = []string{
199+
"web_search",
200+
"web_search_preview",
201+
}
202+
192203
type tokenCounter interface {
193204
Count(model string, input string) (int, error)
194205
}
@@ -527,6 +538,34 @@ func (ce *CostEstimator) EstimateResponseApiTotalCost(model string, usage respon
527538
return inputCost + cachedInputCost + outputCost, err
528539
}
529540

541+
func (ce *CostEstimator) EstimateResponseApiToolCallsCost(tools []responsesOpenai.ToolUnion, model string) (float64, error) {
542+
if len(tools) == 0 {
543+
return 0, nil
544+
}
545+
totalCost := 0.0
546+
for _, tool := range tools {
547+
toolType := tool.Type
548+
cost, ok := OpenAiPerThousandCallsToolCost[extendedToolType(toolType, model)]
549+
if !ok {
550+
return 0, fmt.Errorf("tool type %s is not present in the tool cost map provided", toolType)
551+
}
552+
totalCost += cost
553+
}
554+
return totalCost / 1000, nil
555+
}
556+
557+
var reasoningModelPrefix = []string{"gpt-5", "o1", "o2", "o3"}
558+
559+
func extendedToolType(toolType, model string) string {
560+
if toolType != "web_search_preview" {
561+
return toolType
562+
}
563+
if slices.ContainsFunc(reasoningModelPrefix, func(s string) bool { return strings.HasPrefix(model, s) }) {
564+
return "web_search_preview_reasoning"
565+
}
566+
return toolType
567+
}
568+
530569
func (ce *CostEstimator) estimateResponseApiTokensCost(costMapKey, model string, tks int64) (float64, error) {
531570
costMap, ok := ce.tokenCostMap[costMapKey]
532571
if !ok {

internal/provider/openai/types.go

Lines changed: 30 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -1,31 +1,35 @@
11
package openai
22

33
type ResponseRequest struct {
4-
Background *bool `json:"background,omitzero"`
5-
Conversation *any `json:"conversation,omitzero"`
6-
Include []string `json:"include,omitzero"`
7-
Input *any `json:"input,omitzero"`
8-
Instructions *string `json:"instructions,omitzero"`
9-
MaxOutputTokens *int `json:"max_output_tokens,omitzero"`
10-
MaxToolCalls *int `json:"max_tool_calls,omitzero"`
11-
Metadata *map[string]string `json:"metadata,omitzero"`
12-
Model *string `json:"model,omitzero"`
13-
ParallelToolCalls *bool `json:"parallel_tool_calls,omitzero"`
14-
PreviousResponseId *string `json:"previous_response_id,omitzero"`
15-
Prompt *any `json:"prompt,omitzero"`
16-
PromptCacheKey *string `json:"prompt_cache_key,omitzero"`
17-
Reasoning *any `json:"reasoning,omitzero"`
18-
SafetyIdentifier *string `json:"safety_identifier,omitzero"`
19-
ServiceTier *string `json:"service_tier,omitzero"`
20-
Store *bool `json:"store,omitzero"`
21-
Stream *bool `json:"stream,omitzero"`
22-
StreamOptions *any `json:"stream_options,omitzero"`
23-
Temperature *float32 `json:"temperature,omitzero"`
24-
Text *any `json:"text,omitzero"`
25-
ToolChoice *any `json:"tool_choice,omitzero"`
26-
Tools []any `json:"tools,omitzero"`
27-
TopLogprobs *int `json:"top_logprobs,omitzero"`
28-
TopP *float32 `json:"top_p,omitzero"`
29-
Truncation *string `json:"truncation,omitzero"`
4+
Background *bool `json:"background,omitzero"`
5+
Conversation *any `json:"conversation,omitzero"`
6+
Include []string `json:"include,omitzero"`
7+
Input *any `json:"input,omitzero"`
8+
Instructions *string `json:"instructions,omitzero"`
9+
MaxOutputTokens *int `json:"max_output_tokens,omitzero"`
10+
MaxToolCalls *int `json:"max_tool_calls,omitzero"`
11+
Metadata *map[string]string `json:"metadata,omitzero"`
12+
Model *string `json:"model,omitzero"`
13+
ParallelToolCalls *bool `json:"parallel_tool_calls,omitzero"`
14+
PreviousResponseId *string `json:"previous_response_id,omitzero"`
15+
Prompt *any `json:"prompt,omitzero"`
16+
PromptCacheKey *string `json:"prompt_cache_key,omitzero"`
17+
Reasoning *any `json:"reasoning,omitzero"`
18+
SafetyIdentifier *string `json:"safety_identifier,omitzero"`
19+
ServiceTier *string `json:"service_tier,omitzero"`
20+
Store *bool `json:"store,omitzero"`
21+
Stream *bool `json:"stream,omitzero"`
22+
StreamOptions *any `json:"stream_options,omitzero"`
23+
Temperature *float32 `json:"temperature,omitzero"`
24+
Text *any `json:"text,omitzero"`
25+
ToolChoice *any `json:"tool_choice,omitzero"`
26+
Tools []ResponseRequestToolUnion `json:"tools,omitzero"`
27+
TopLogprobs *int `json:"top_logprobs,omitzero"`
28+
TopP *float32 `json:"top_p,omitzero"`
29+
Truncation *string `json:"truncation,omitzero"`
3030
//User *string `json:"user,omitzero"` //Deprecated
3131
}
32+
33+
type ResponseRequestToolUnion struct {
34+
Type string `json:"type"`
35+
}

internal/server/web/proxy/middleware.go

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@ import (
66
"fmt"
77
"io"
88
"net/http"
9+
"slices"
910
"strconv"
1011
"strings"
1112
"time"
@@ -57,6 +58,7 @@ type estimator interface {
5758
EstimateEmbeddingsInputCost(model string, tks int) (float64, error)
5859
EstimateChatCompletionPromptTokenCounts(model string, r *goopenai.ChatCompletionRequest) (int, error)
5960
EstimateResponseApiTotalCost(model string, usage responsesOpenai.ResponseUsage) (float64, error)
61+
EstimateResponseApiToolCallsCost(tools []responsesOpenai.ToolUnion, model string) (float64, error)
6062
}
6163

6264
type azureEstimator interface {
@@ -812,6 +814,21 @@ func getMiddleware(cpm CustomProvidersManager, rm routeManager, pm PoliciesManag
812814
return
813815
}
814816

817+
hasNotAllowedTools := false
818+
for _, tool := range responsesReq.Tools {
819+
if !slices.Contains(openai.AllowedTools, tool.Type) {
820+
hasNotAllowedTools = true
821+
break
822+
}
823+
}
824+
825+
if hasNotAllowedTools {
826+
telemetry.Incr("bricksllm.proxy.get_middleware.tool_not_allowed", nil, 1)
827+
JSON(c, http.StatusForbidden, "[BricksLLM] one of the tools is not allowed")
828+
c.Abort()
829+
return
830+
}
831+
815832
userId = gopointer.ToValueOrDefault(responsesReq.SafetyIdentifier, "")
816833
enrichedEvent.Request = responsesReq
817834
c.Set("model", gopointer.ToValueOrDefault(responsesReq.Model, ""))

internal/server/web/proxy/responses.go

Lines changed: 12 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -98,6 +98,12 @@ func getResponsesHandler(prod, private bool, client http.Client, e estimator) gi
9898
telemetry.Incr("bricksllm.proxy.get_chat_completion_handler.estimate_total_cost_error", nil, 1)
9999
logError(log, "error when estimating openai cost", prod, err)
100100
}
101+
toolsCost, err := e.EstimateResponseApiToolCallsCost(resp.Tools, model)
102+
if err != nil {
103+
telemetry.Incr("bricksllm.proxy.get_chat_completion_handler.estimate_tool_calls_cost_error", nil, 1)
104+
logError(log, "error when estimating openai tool calls cost", prod, err)
105+
}
106+
cost += toolsCost
101107
}
102108

103109
c.Set("costInUsd", cost)
@@ -231,6 +237,12 @@ func getResponsesHandler(prod, private bool, client http.Client, e estimator) gi
231237
telemetry.Incr("bricksllm.proxy.get_chat_completion_handler.estimate_total_cost_error", nil, 1)
232238
logError(log, "error when estimating openai cost", prod, err)
233239
}
240+
toolsCost, err := e.EstimateResponseApiToolCallsCost(responsesStreamResp.Response.Tools, model)
241+
if err != nil {
242+
telemetry.Incr("bricksllm.proxy.get_chat_completion_handler.estimate_tool_calls_cost_error", nil, 1)
243+
logError(log, "error when estimating openai tool calls cost", prod, err)
244+
}
245+
streamCost += toolsCost
234246
streamPromptTokenCount, err = int64ToInt(responsesStreamResp.Response.Usage.InputTokens)
235247
if err != nil {
236248
telemetry.Incr("bricksllm.proxy.get_responses_handler.int64_to_int_error", nil, 1)

0 commit comments

Comments
 (0)