148 lines
3.7 KiB
Go
148 lines
3.7 KiB
Go
package proxy
|
|
|
|
import (
|
|
"bytes"
|
|
"encoding/json"
|
|
"io"
|
|
"sync"
|
|
)
|
|
|
|
// maxParseBody bounds the non-streamed JSON body accumulated to extract usage/timings: beyond it we
|
|
// record no tokens rather than hold an unbounded body in memory.
|
|
const maxParseBody = 1 << 20 // 1 MiB
|
|
|
|
// usageTimings is the last usage/timings object seen in a stream, or the single object parsed from a
|
|
// non-streamed JSON body.
|
|
type usageTimings struct {
|
|
hasUsage bool
|
|
usage struct{ prompt, completion int }
|
|
hasTimings bool
|
|
timings struct{ promptN, cacheN, predictedN int }
|
|
}
|
|
|
|
// tokens resolves the recorded usage and timings to the three counts the accounting row needs.
|
|
// prompt and completion come from usage when present, else from timings; cached comes only from
|
|
// timings.
|
|
func (u *usageTimings) tokens() (prompt, cached, completion int) {
|
|
if u.hasUsage {
|
|
prompt = u.usage.prompt
|
|
completion = u.usage.completion
|
|
} else if u.hasTimings {
|
|
prompt = u.timings.promptN
|
|
completion = u.timings.predictedN
|
|
}
|
|
if u.hasTimings {
|
|
cached = u.timings.cacheN
|
|
}
|
|
return
|
|
}
|
|
|
|
// tee wraps a response body, passing every byte through unchanged while scanning for usage/timings.
|
|
// For a text/event-stream body it scans complete data: lines and remembers the last object seen; for
|
|
// anything else it accumulates the body (bounded) and parses it once at Close.
|
|
type tee struct {
|
|
body io.Reader
|
|
closed bool
|
|
streamed bool
|
|
pending []byte
|
|
buf bytes.Buffer
|
|
mu sync.Mutex
|
|
last usageTimings
|
|
}
|
|
|
|
func newTee(r io.Reader, streamed bool) *tee {
|
|
return &tee{body: r, streamed: streamed}
|
|
}
|
|
|
|
// Read reads from the upstream body and feeds the bytes to the scanner without holding any back.
|
|
func (t *tee) Read(b []byte) (int, error) {
|
|
n, err := t.body.Read(b)
|
|
if n > 0 {
|
|
t.ingest(b[:n])
|
|
}
|
|
return n, err
|
|
}
|
|
|
|
// ingest routes a fresh chunk to the streamed scanner or the non-streamed accumulator.
|
|
func (t *tee) ingest(p []byte) {
|
|
t.mu.Lock()
|
|
defer t.mu.Unlock()
|
|
if t.streamed {
|
|
t.scanSSE(p)
|
|
return
|
|
}
|
|
if t.buf.Len() < maxParseBody {
|
|
t.buf.Write(p)
|
|
}
|
|
}
|
|
|
|
// scanSSE splits complete lines off the pending buffer and parses each "data: " line.
|
|
func (t *tee) scanSSE(p []byte) {
|
|
t.pending = append(t.pending, p...)
|
|
for {
|
|
i := bytes.IndexByte(t.pending, '\n')
|
|
if i < 0 {
|
|
return
|
|
}
|
|
line := t.pending[:i]
|
|
t.pending = t.pending[i+1:]
|
|
if bytes.HasPrefix(line, []byte("data: ")) {
|
|
t.parseData(line[len("data: "):])
|
|
}
|
|
}
|
|
}
|
|
|
|
// parseData decodes one data line's JSON object and remembers its usage and/or timings.
|
|
func (t *tee) parseData(data []byte) {
|
|
var doc struct {
|
|
Usage *struct {
|
|
PromptTokens int `json:"prompt_tokens"`
|
|
CompletionTokens int `json:"completion_tokens"`
|
|
} `json:"usage"`
|
|
Timings *struct {
|
|
PromptN int `json:"prompt_n"`
|
|
CacheN int `json:"cache_n"`
|
|
PredictedN int `json:"predicted_n"`
|
|
} `json:"timings"`
|
|
}
|
|
if err := json.Unmarshal(data, &doc); err != nil {
|
|
return
|
|
}
|
|
if doc.Usage != nil {
|
|
t.last.hasUsage = true
|
|
t.last.usage.prompt = doc.Usage.PromptTokens
|
|
t.last.usage.completion = doc.Usage.CompletionTokens
|
|
}
|
|
if doc.Timings != nil {
|
|
t.last.hasTimings = true
|
|
t.last.timings.promptN = doc.Timings.PromptN
|
|
t.last.timings.cacheN = doc.Timings.CacheN
|
|
t.last.timings.predictedN = doc.Timings.PredictedN
|
|
}
|
|
}
|
|
|
|
// Close closes the underlying body, parsing a non-streamed JSON body once at the end.
|
|
func (t *tee) Close() error {
|
|
t.mu.Lock()
|
|
if t.closed {
|
|
t.mu.Unlock()
|
|
return nil
|
|
}
|
|
t.closed = true
|
|
if !t.streamed && t.buf.Len() > 0 {
|
|
t.parseData(t.buf.Bytes())
|
|
}
|
|
t.mu.Unlock()
|
|
if c, ok := t.body.(io.Closer); ok {
|
|
return c.Close()
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// tokens returns the last usage/timings seen, if any.
|
|
func (t *tee) tokens() (prompt, cached, completion int) {
|
|
t.mu.Lock()
|
|
defer t.mu.Unlock()
|
|
return t.last.tokens()
|
|
}
|