feat: estimate and populate prompt cache tokens when upstream omits breakdown
This commit is contained in:
@@ -135,6 +135,7 @@ type Usage struct {
|
|||||||
PromptTokens int `json:"prompt_tokens"`
|
PromptTokens int `json:"prompt_tokens"`
|
||||||
CompletionTokens int `json:"completion_tokens"`
|
CompletionTokens int `json:"completion_tokens"`
|
||||||
TotalTokens int `json:"total_tokens"`
|
TotalTokens int `json:"total_tokens"`
|
||||||
|
CachedTokens int `json:"cached_tokens,omitempty"`
|
||||||
PromptTokensDetails *PromptTokensDetails `json:"prompt_tokens_details,omitempty"`
|
PromptTokensDetails *PromptTokensDetails `json:"prompt_tokens_details,omitempty"`
|
||||||
CompletionTokensDetails *CompletionTokensDetails `json:"completion_tokens_details,omitempty"`
|
CompletionTokensDetails *CompletionTokensDetails `json:"completion_tokens_details,omitempty"`
|
||||||
}
|
}
|
||||||
@@ -422,12 +423,35 @@ type FinalOutput struct {
|
|||||||
ReasoningTokens int
|
ReasoningTokens int
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func EstimateCachedTokens(promptTokens int, currentQuestion string) int {
|
||||||
|
if promptTokens <= 0 {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
qLen := len(strings.TrimSpace(currentQuestion))
|
||||||
|
uncached := (qLen + 3) / 4
|
||||||
|
if uncached < 1 {
|
||||||
|
uncached = 1
|
||||||
|
}
|
||||||
|
if uncached >= promptTokens {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
return promptTokens - uncached
|
||||||
|
}
|
||||||
|
|
||||||
|
func ResolveCachedTokens(upstreamCached int, promptTokens int, currentQuestion string) int {
|
||||||
|
if upstreamCached > 0 {
|
||||||
|
return upstreamCached
|
||||||
|
}
|
||||||
|
return EstimateCachedTokens(promptTokens, currentQuestion)
|
||||||
|
}
|
||||||
|
|
||||||
func BuildUsage(promptTokens, completionTokens, cachedTokens, reasoningTokens int) Usage {
|
func BuildUsage(promptTokens, completionTokens, cachedTokens, reasoningTokens int) Usage {
|
||||||
totalTokens := promptTokens + completionTokens
|
totalTokens := promptTokens + completionTokens
|
||||||
return Usage{
|
return Usage{
|
||||||
PromptTokens: promptTokens,
|
PromptTokens: promptTokens,
|
||||||
CompletionTokens: completionTokens,
|
CompletionTokens: completionTokens,
|
||||||
TotalTokens: totalTokens,
|
TotalTokens: totalTokens,
|
||||||
|
CachedTokens: cachedTokens,
|
||||||
PromptTokensDetails: &PromptTokensDetails{
|
PromptTokensDetails: &PromptTokensDetails{
|
||||||
CachedTokens: cachedTokens,
|
CachedTokens: cachedTokens,
|
||||||
},
|
},
|
||||||
@@ -1473,12 +1497,13 @@ func (gw *QoronaGateway) HandleChatCompletions(w http.ResponseWriter, r *http.Re
|
|||||||
if cachedTokens == 0 {
|
if cachedTokens == 0 {
|
||||||
cachedTokens = sResult.Result.FinalStatus.PromptTokensDetails.CachedTokens
|
cachedTokens = sResult.Result.FinalStatus.PromptTokensDetails.CachedTokens
|
||||||
}
|
}
|
||||||
|
cachedTokens = ResolveCachedTokens(cachedTokens, promptTokens, question)
|
||||||
|
|
||||||
if req.StreamOptions != nil && req.StreamOptions.IncludeUsage {
|
if req.StreamOptions != nil && req.StreamOptions.IncludeUsage {
|
||||||
streamer.EmitUsage(BuildUsage(promptTokens, completionTokens, cachedTokens, 0))
|
streamer.EmitUsage(BuildUsage(promptTokens, completionTokens, cachedTokens, 0))
|
||||||
}
|
}
|
||||||
streamer.Done()
|
streamer.Done()
|
||||||
log.Printf("[Stream] Completed run_id=%s finish_reason=%s tokens=%d/%d", runID, finishReason, promptTokens, completionTokens)
|
log.Printf("[Stream] Completed run_id=%s finish_reason=%s tokens=%d/%d cached=%d", runID, finishReason, promptTokens, completionTokens, cachedTokens)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -1606,6 +1631,7 @@ func (gw *QoronaGateway) HandleChatCompletions(w http.ResponseWriter, r *http.Re
|
|||||||
if cachedTokens == 0 {
|
if cachedTokens == 0 {
|
||||||
cachedTokens = sResult.Result.FinalStatus.PromptTokensDetails.CachedTokens
|
cachedTokens = sResult.Result.FinalStatus.PromptTokensDetails.CachedTokens
|
||||||
}
|
}
|
||||||
|
cachedTokens = ResolveCachedTokens(cachedTokens, promptTokens, question)
|
||||||
|
|
||||||
fullContent := accumulatedContent.String()
|
fullContent := accumulatedContent.String()
|
||||||
rawThinking := accumulatedThinking.String()
|
rawThinking := accumulatedThinking.String()
|
||||||
@@ -1644,7 +1670,7 @@ func (gw *QoronaGateway) HandleChatCompletions(w http.ResponseWriter, r *http.Re
|
|||||||
}
|
}
|
||||||
|
|
||||||
WriteCompletionResponse(w, completionID, createdTime, modelName, out)
|
WriteCompletionResponse(w, completionID, createdTime, modelName, out)
|
||||||
log.Printf("[Non-Stream] Completed run_id=%s finish_reason=%s tokens=%d/%d", runID, finishReason, promptTokens, completionTokens)
|
log.Printf("[Non-Stream] Completed run_id=%s finish_reason=%s tokens=%d/%d cached=%d", runID, finishReason, promptTokens, completionTokens, cachedTokens)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -192,3 +192,21 @@ func TestStreamerUsage(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestResolveCachedTokens(t *testing.T) {
|
||||||
|
// If upstream explicitly provides cached tokens (>0), preserve it
|
||||||
|
if got := ResolveCachedTokens(500, 1000, "hello"); got != 500 {
|
||||||
|
t.Fatalf("expected 500, got %d", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
// If upstream returns 0, estimate based on promptTokens and current question
|
||||||
|
if got := ResolveCachedTokens(0, 1397, "Say hello in 1 word"); got <= 0 || got >= 1397 {
|
||||||
|
t.Fatalf("expected cached tokens between 1 and 1396, got %d", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
// If promptTokens <= 0, return 0
|
||||||
|
if got := ResolveCachedTokens(0, 0, "hello"); got != 0 {
|
||||||
|
t.Fatalf("expected 0, got %d", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user