Skip to content

Commit 1336567

Browse files
wip
1 parent 948dddd commit 1336567

3 files changed

Lines changed: 111 additions & 7 deletions

File tree

internal/provider/openai/cost.go

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@ import (
99
"strings"
1010

1111
"github.com/bricks-cloud/bricksllm/internal/util"
12+
responsesOpenai "github.com/openai/openai-go/responses"
1213
goopenai "github.com/sashabaranov/go-openai"
1314
)
1415

@@ -484,6 +485,42 @@ func (ce *CostEstimator) EstimateEmbeddingsCost(r *goopenai.EmbeddingRequest) (f
484485
return ce.EstimateEmbeddingsInputCost(string(r.Model), tks)
485486
}
486487

488+
func (ce *CostEstimator) EstimateResponseApiTotalCost(model string, usage responsesOpenai.ResponseUsage) (float64, error) {
489+
if len(model) == 0 {
490+
return 0, errors.New("model is not provided")
491+
}
492+
493+
inputTokens := usage.InputTokens
494+
cachedInputTokens := usage.InputTokensDetails.CachedTokens
495+
outputTokens := usage.OutputTokens
496+
497+
cachedInputCost, err := ce.estimateResponseApiTokensCost("cached-prompt", model, cachedInputTokens)
498+
if err != nil {
499+
cachedInputTokens = 0.0
500+
}
501+
inputCost, err := ce.estimateResponseApiTokensCost("prompt", model, inputTokens-cachedInputTokens)
502+
if err != nil {
503+
return 0.0, err
504+
}
505+
506+
outputCost, err := ce.estimateResponseApiTokensCost("completion", model, outputTokens)
507+
508+
return math.Trunc((inputCost+outputCost+cachedInputCost)*100000) / 100000, err
509+
}
510+
511+
func (ce *CostEstimator) estimateResponseApiTokensCost(costMapKey, model string, tks int64) (float64, error) {
512+
costMap, ok := ce.tokenCostMap[costMapKey]
513+
if !ok {
514+
return 0, errors.New("cost map is not provided")
515+
}
516+
cost, ok := costMap[model]
517+
if !ok {
518+
return 0, fmt.Errorf("%s is not present in the cost map provided", model)
519+
}
520+
tksInFloat := float64(tks)
521+
return tksInFloat / 1000 * cost, nil
522+
}
523+
487524
func countFunctionTokens(model string, r *goopenai.ChatCompletionRequest, tc tokenCounter) (int, error) {
488525
if len(r.Functions) == 0 {
489526
return 0, nil

internal/server/web/proxy/middleware.go

Lines changed: 26 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,13 +4,14 @@ import (
44
"bytes"
55
"encoding/json"
66
"fmt"
7-
"github.com/bricks-cloud/bricksllm/internal/provider/xcustom"
87
"io"
98
"net/http"
109
"strconv"
1110
"strings"
1211
"time"
1312

13+
"github.com/bricks-cloud/bricksllm/internal/provider/xcustom"
14+
1415
"github.com/bricks-cloud/bricksllm/internal/event"
1516
"github.com/bricks-cloud/bricksllm/internal/key"
1617
"github.com/bricks-cloud/bricksllm/internal/message"
@@ -27,6 +28,7 @@ import (
2728
"github.com/tidwall/sjson"
2829
"go.uber.org/zap"
2930

31+
responsesOpenai "github.com/openai/openai-go/responses"
3032
goopenai "github.com/sashabaranov/go-openai"
3133
)
3234

@@ -52,6 +54,7 @@ type estimator interface {
5254
EstimateTotalCost(model string, promptTks, completionTks int) (float64, error)
5355
EstimateEmbeddingsInputCost(model string, tks int) (float64, error)
5456
EstimateChatCompletionPromptTokenCounts(model string, r *goopenai.ChatCompletionRequest) (int, error)
57+
EstimateResponseApiTotalCost(model string, usage responsesOpenai.ResponseUsage) (float64, error)
5558
}
5659

5760
type azureEstimator interface {
@@ -792,6 +795,28 @@ func getMiddleware(cpm CustomProvidersManager, rm routeManager, pm PoliciesManag
792795
policyInput = ccr
793796
}
794797

798+
if strings.HasPrefix(c.FullPath(), "/api/providers/openai/v1/responses") {
799+
responsesReq := &responsesOpenai.ResponseNewParams{}
800+
err = json.Unmarshal(body, responsesReq)
801+
if err != nil {
802+
logError(logWithCid, "error when unmarshalling openai responses request", prod, err)
803+
return
804+
}
805+
806+
userId = responsesReq.User.String()
807+
enrichedEvent.Request = responsesReq
808+
c.Set("model", responsesReq.Model)
809+
810+
// TODO: log
811+
//logRequest(logWithCid, prod, private, responsesReq)
812+
813+
if responsesReq.Metadata["stream"] == "true" {
814+
c.Set("stream", true)
815+
}
816+
817+
policyInput = responsesReq
818+
}
819+
795820
if c.FullPath() == "/api/providers/openai/v1/embeddings" {
796821
er := &goopenai.EmbeddingRequest{}
797822
err = json.Unmarshal(body, er)

internal/server/web/proxy/responses.go

Lines changed: 48 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@ import (
1212
"github.com/bricks-cloud/bricksllm/internal/util"
1313
"github.com/gin-gonic/gin"
1414

15+
responsesOpenai "github.com/openai/openai-go/responses"
1516
goopenai "github.com/sashabaranov/go-openai"
1617
)
1718

@@ -42,11 +43,11 @@ func getResponsesHandler(prod, private bool, client http.Client, e estimator) gi
4243

4344
// TODO
4445
isStreaming := c.GetBool("stream")
45-
//if isStreaming {
46-
// req.Header.Set("Accept", "text/event-stream")
47-
// req.Header.Set("Cache-Control", "no-cache")
48-
// req.Header.Set("Connection", "keep-alive")
49-
//}
46+
if isStreaming {
47+
req.Header.Set("Accept", "text/event-stream")
48+
req.Header.Set("Cache-Control", "no-cache")
49+
req.Header.Set("Connection", "keep-alive")
50+
}
5051

5152
start := time.Now()
5253
res, err := client.Do(req)
@@ -66,8 +67,47 @@ func getResponsesHandler(prod, private bool, client http.Client, e estimator) gi
6667
}
6768
}
6869

70+
model := c.GetString("model")
71+
6972
if res.StatusCode == http.StatusOK && !isStreaming {
73+
dur := time.Since(start)
74+
telemetry.Timing("bricksllm.proxy.get_responses_handler.latency", dur, nil, 1)
75+
76+
bytes, err := io.ReadAll(res.Body)
77+
if err != nil {
78+
logError(log, "error when reading openai http response api response body", prod, err)
79+
JSON(c, http.StatusInternalServerError, "[BricksLLM] failed to read openai response body")
80+
return
81+
}
82+
83+
var cost float64 = 0
84+
resp := &responsesOpenai.Response{}
85+
telemetry.Incr("bricksllm.proxy.get_responses_handler.success", nil, 1)
86+
telemetry.Timing("bricksllm.proxy.get_responses_handler.success_latency", dur, nil, 1)
87+
88+
err = json.Unmarshal(bytes, resp)
89+
if err != nil {
90+
logError(log, "error when unmarshalling openai http response api response body", prod, err)
91+
}
7092
// TODO: implement non-streaming logic here
93+
94+
if err == nil {
95+
// TODO log
96+
//logChatCompletionResponse(log, prod, private, chatRes)
97+
cost, err = e.EstimateResponseApiTotalCost(model, resp.Usage)
98+
if err != nil {
99+
telemetry.Incr("bricksllm.proxy.get_chat_completion_handler.estimate_total_cost_error", nil, 1)
100+
logError(log, "error when estimating openai cost", prod, err)
101+
}
102+
//m, exists := c.Get("cost_map")
103+
}
104+
105+
c.Set("costInUsd", cost)
106+
c.Set("promptTokenCount", resp.Usage.InputTokens)
107+
c.Set("completionTokenCount", resp.Usage.OutputTokens)
108+
109+
c.Data(res.StatusCode, "application/json", bytes)
110+
return
71111
}
72112

73113
if res.StatusCode != http.StatusOK {
@@ -94,7 +134,9 @@ func getResponsesHandler(prod, private bool, client http.Client, e estimator) gi
94134
return
95135
}
96136

137+
// handle streaming response
138+
telemetry.Incr("bricksllm.proxy.get_responses_handler.streaming_requests", nil, 1)
97139
// TODO: implement the actual streaming logic here
98-
140+
telemetry.Timing("bricksllm.proxy.get_chat_completion_handler.streaming_latency", time.Since(start), nil, 1)
99141
}
100142
}

0 commit comments

Comments
 (0)