mirror of
https://github.com/ollama/ollama.git
synced 2026-09-21 13:38:14 -05:00
Report cached prompt tokens (#17943)
* Report cached prompt tokens Add prompt_eval_cached_count to native responses and expose equivalent cached-token fields through the OpenAI- and Anthropic-compatible APIs. Keep prompt_eval_count as the logical input total while excluding cache hits from CLI and benchmark prefill rates. Surface processed and cached prompt counts in benchmark output. Collect cache counts from llama-server and MLX, preserve coherent metrics across two-pass structured generation. Fixes #8008 Related to #15758 * review comments
This commit is contained in:
+44
-16
@@ -217,8 +217,31 @@ type MessagesResponse struct {
|
||||
|
||||
// Usage contains token usage information
|
||||
type Usage struct {
|
||||
InputTokens int `json:"input_tokens"`
|
||||
OutputTokens int `json:"output_tokens"`
|
||||
InputTokens int `json:"input_tokens"`
|
||||
CacheReadInputTokens *int `json:"cache_read_input_tokens,omitempty"`
|
||||
OutputTokens int `json:"output_tokens"`
|
||||
}
|
||||
|
||||
// UsageFromMetrics separates total prompt tokens into uncached and cache-read counts.
|
||||
func UsageFromMetrics(metrics api.Metrics) Usage {
|
||||
total := max(0, metrics.PromptEvalCount)
|
||||
var cached *int
|
||||
if metrics.PromptEvalCachedCount != nil {
|
||||
count := min(max(0, *metrics.PromptEvalCachedCount), total)
|
||||
cached = &count
|
||||
}
|
||||
return Usage{
|
||||
InputTokens: total - intValue(cached),
|
||||
CacheReadInputTokens: cached,
|
||||
OutputTokens: metrics.EvalCount,
|
||||
}
|
||||
}
|
||||
|
||||
func intValue(v *int) int {
|
||||
if v == nil {
|
||||
return 0
|
||||
}
|
||||
return *v
|
||||
}
|
||||
|
||||
// Streaming event types
|
||||
@@ -273,8 +296,9 @@ type MessageDelta struct {
|
||||
|
||||
// DeltaUsage contains cumulative token usage
|
||||
type DeltaUsage struct {
|
||||
InputTokens int `json:"input_tokens"`
|
||||
OutputTokens int `json:"output_tokens"`
|
||||
InputTokens int `json:"input_tokens"`
|
||||
CacheReadInputTokens *int `json:"cache_read_input_tokens,omitempty"`
|
||||
OutputTokens int `json:"output_tokens"`
|
||||
}
|
||||
|
||||
// MessageStopEvent signals the end of the message
|
||||
@@ -688,10 +712,7 @@ func ToMessagesResponse(id string, r api.ChatResponse) MessagesResponse {
|
||||
Model: r.Model,
|
||||
Content: content,
|
||||
StopReason: stopReason,
|
||||
Usage: Usage{
|
||||
InputTokens: r.Metrics.PromptEvalCount,
|
||||
OutputTokens: r.Metrics.EvalCount,
|
||||
},
|
||||
Usage: UsageFromMetrics(r.Metrics),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -721,6 +742,7 @@ type StreamConverter struct {
|
||||
firstWrite bool
|
||||
contentIndex int
|
||||
inputTokens int
|
||||
cacheReadTokens *int
|
||||
outputTokens int
|
||||
estimatedInputTokens int // Estimated tokens from request (used when actual metrics are 0)
|
||||
thinkingStarted bool
|
||||
@@ -752,8 +774,10 @@ func (c *StreamConverter) Process(r api.ChatResponse) []StreamEvent {
|
||||
if c.firstWrite {
|
||||
c.firstWrite = false
|
||||
// Use actual metrics if available, otherwise use estimate
|
||||
c.inputTokens = r.Metrics.PromptEvalCount
|
||||
if c.inputTokens == 0 && c.estimatedInputTokens > 0 {
|
||||
usage := UsageFromMetrics(r.Metrics)
|
||||
c.inputTokens = usage.InputTokens
|
||||
c.cacheReadTokens = usage.CacheReadInputTokens
|
||||
if c.inputTokens == 0 && intValue(c.cacheReadTokens) == 0 && c.estimatedInputTokens > 0 {
|
||||
c.inputTokens = c.estimatedInputTokens
|
||||
}
|
||||
|
||||
@@ -768,8 +792,9 @@ func (c *StreamConverter) Process(r api.ChatResponse) []StreamEvent {
|
||||
Model: c.Model,
|
||||
Content: []ContentBlock{},
|
||||
Usage: Usage{
|
||||
InputTokens: c.inputTokens,
|
||||
OutputTokens: 0,
|
||||
InputTokens: c.inputTokens,
|
||||
CacheReadInputTokens: c.cacheReadTokens,
|
||||
OutputTokens: 0,
|
||||
},
|
||||
},
|
||||
},
|
||||
@@ -950,8 +975,10 @@ func (c *StreamConverter) Process(r api.ChatResponse) []StreamEvent {
|
||||
})
|
||||
}
|
||||
|
||||
c.inputTokens = r.Metrics.PromptEvalCount
|
||||
c.outputTokens = r.Metrics.EvalCount
|
||||
usage := UsageFromMetrics(r.Metrics)
|
||||
c.inputTokens = usage.InputTokens
|
||||
c.cacheReadTokens = usage.CacheReadInputTokens
|
||||
c.outputTokens = usage.OutputTokens
|
||||
stopReason := mapStopReason(r.DoneReason, len(c.toolCallsSent) > 0)
|
||||
|
||||
events = append(events, StreamEvent{
|
||||
@@ -962,8 +989,9 @@ func (c *StreamConverter) Process(r api.ChatResponse) []StreamEvent {
|
||||
StopReason: stopReason,
|
||||
},
|
||||
Usage: DeltaUsage{
|
||||
InputTokens: c.inputTokens,
|
||||
OutputTokens: c.outputTokens,
|
||||
InputTokens: c.inputTokens,
|
||||
CacheReadInputTokens: c.cacheReadTokens,
|
||||
OutputTokens: c.outputTokens,
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
@@ -16,6 +16,10 @@ const (
|
||||
testImage = `iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk+A8AAQUBAScY42YAAAAASUVORK5CYII=`
|
||||
)
|
||||
|
||||
func testIntPtr(v int) *int {
|
||||
return &v
|
||||
}
|
||||
|
||||
// textContent is a convenience for constructing []ContentBlock with a single text block in tests.
|
||||
func textContent(s string) []ContentBlock {
|
||||
return []ContentBlock{{Type: "text", Text: &s}}
|
||||
@@ -30,6 +34,61 @@ func makeArgs(kvs ...any) api.ToolCallFunctionArguments {
|
||||
return args
|
||||
}
|
||||
|
||||
func TestUsageFromMetricsBoundsCacheReads(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
metrics api.Metrics
|
||||
want Usage
|
||||
}{
|
||||
{
|
||||
name: "negative counts",
|
||||
metrics: api.Metrics{PromptEvalCount: -1, PromptEvalCachedCount: testIntPtr(-2), EvalCount: 3},
|
||||
want: Usage{CacheReadInputTokens: testIntPtr(0), OutputTokens: 3},
|
||||
},
|
||||
{
|
||||
name: "cache reads exceed prompt",
|
||||
metrics: api.Metrics{PromptEvalCount: 3, PromptEvalCachedCount: testIntPtr(5), EvalCount: 2},
|
||||
want: Usage{CacheReadInputTokens: testIntPtr(3), OutputTokens: 2},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if diff := cmp.Diff(tt.want, UsageFromMetrics(tt.metrics)); diff != "" {
|
||||
t.Errorf("usage mismatch (-want +got):\n%s", diff)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestUsageCacheReadJSON(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
count *int
|
||||
want string
|
||||
}{
|
||||
{name: "unreported", want: `{"input_tokens":10,"output_tokens":2}`},
|
||||
{name: "zero", count: testIntPtr(0), want: `{"input_tokens":10,"cache_read_input_tokens":0,"output_tokens":2}`},
|
||||
{name: "positive", count: testIntPtr(4), want: `{"input_tokens":6,"cache_read_input_tokens":4,"output_tokens":2}`},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
data, err := json.Marshal(UsageFromMetrics(api.Metrics{
|
||||
PromptEvalCount: 10,
|
||||
PromptEvalCachedCount: tt.count,
|
||||
EvalCount: 2,
|
||||
}))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := string(data); got != tt.want {
|
||||
t.Errorf("json = %s, want %s", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestFromMessagesRequest_Basic(t *testing.T) {
|
||||
req := MessagesRequest{
|
||||
Model: "test-model",
|
||||
@@ -861,8 +920,9 @@ func TestToMessagesResponse_Basic(t *testing.T) {
|
||||
Done: true,
|
||||
DoneReason: "stop",
|
||||
Metrics: api.Metrics{
|
||||
PromptEvalCount: 10,
|
||||
EvalCount: 5,
|
||||
PromptEvalCount: 10,
|
||||
PromptEvalCachedCount: testIntPtr(4),
|
||||
EvalCount: 5,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -886,9 +946,17 @@ func TestToMessagesResponse_Basic(t *testing.T) {
|
||||
if result.StopReason != "end_turn" {
|
||||
t.Errorf("expected stop_reason 'end_turn', got %q", result.StopReason)
|
||||
}
|
||||
if result.Usage.InputTokens != 10 || result.Usage.OutputTokens != 5 {
|
||||
if result.Usage.InputTokens != 6 || intValue(result.Usage.CacheReadInputTokens) != 4 || result.Usage.OutputTokens != 5 {
|
||||
t.Errorf("unexpected usage: %+v", result.Usage)
|
||||
}
|
||||
|
||||
data, err := json.Marshal(result.Usage)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(string(data), `"cache_read_input_tokens":4`) {
|
||||
t.Errorf("unexpected usage json: %s", data)
|
||||
}
|
||||
}
|
||||
|
||||
func TestToMessagesResponse_PreservesClaudeAutoClassifierOutput(t *testing.T) {
|
||||
@@ -1072,7 +1140,7 @@ func TestStreamConverter_Basic(t *testing.T) {
|
||||
Role: "assistant",
|
||||
Content: "Hello",
|
||||
},
|
||||
Metrics: api.Metrics{PromptEvalCount: 10},
|
||||
Metrics: api.Metrics{PromptEvalCount: 10, PromptEvalCachedCount: testIntPtr(4)},
|
||||
}
|
||||
|
||||
events1 := conv.Process(resp1)
|
||||
@@ -1100,7 +1168,7 @@ func TestStreamConverter_Basic(t *testing.T) {
|
||||
},
|
||||
Done: true,
|
||||
DoneReason: "stop",
|
||||
Metrics: api.Metrics{PromptEvalCount: 10, EvalCount: 5},
|
||||
Metrics: api.Metrics{PromptEvalCount: 10, PromptEvalCachedCount: testIntPtr(4), EvalCount: 5},
|
||||
}
|
||||
|
||||
events2 := conv.Process(resp2)
|
||||
@@ -1118,7 +1186,7 @@ func TestStreamConverter_Basic(t *testing.T) {
|
||||
t.Errorf("unexpected stop reason: %+v", data.Delta.StopReason)
|
||||
}
|
||||
|
||||
if data.Usage.InputTokens != 10 || data.Usage.OutputTokens != 5 {
|
||||
if data.Usage.InputTokens != 6 || intValue(data.Usage.CacheReadInputTokens) != 4 || data.Usage.OutputTokens != 5 {
|
||||
t.Errorf("unexpected usage: %+v", data.Usage)
|
||||
}
|
||||
} else {
|
||||
|
||||
Reference in New Issue
Block a user