Files
ResolveSpec/pkg/aiproxy/usage.go
T
warkanum 63cb0d22eb feat(aiproxy): add authenticated proxy for OpenAI-compatible APIs and MCP servers
Hides upstream URLs and keys behind the security layer. Per-upstream roles,
model/tool allowlists, per-user rate limits, Before/AfterProxy hooks, audit
sink, Prometheus metrics (metrics.AIProxyRecorder) and upstreams loaded from
the resolvespec_ai_proxies stored procedure.
2026-10-07 23:39:40 +02:00

119 lines
2.4 KiB
Go

package aiproxy
import (
"bytes"
"io"
"sync"
"github.com/tidwall/gjson"
)
const maxCapture = 1 << 20
// watchBody passes the response through while it looks for OpenAI usage, and fires
// AfterProxy once the body ends or is closed.
type watchBody struct {
rc io.ReadCloser
capture bool // read usage
sse bool
status int
c *call
buf bytes.Buffer // JSON body (capped)
overflow bool
line []byte // partial SSE line
usage Usage
once sync.Once
}
func newWatchBody(rc io.ReadCloser, openai, sse bool, status int, c *call) io.ReadCloser {
return &watchBody{rc: rc, capture: openai, sse: sse, status: status, c: c}
}
func (w *watchBody) Read(p []byte) (int, error) {
n, err := w.rc.Read(p)
if w.capture && n > 0 {
w.feed(p[:n])
}
if err != nil {
w.end(err)
}
return n, err
}
func (w *watchBody) Close() error {
err := w.rc.Close()
w.end(nil)
return err
}
func (w *watchBody) end(err error) {
w.once.Do(func() {
if w.capture && !w.sse && !w.overflow {
if u, ok := parseUsage(w.buf.Bytes()); ok {
w.usage = u
}
}
if err == io.EOF {
err = nil
}
w.c.finish(w.status, w.usage, err)
})
}
func (w *watchBody) feed(b []byte) {
if !w.sse {
if w.overflow || w.buf.Len()+len(b) > maxCapture {
w.overflow = true
return
}
w.buf.Write(b)
return
}
w.line = append(w.line, b...)
for {
i := bytes.IndexByte(w.line, '\n')
if i < 0 {
break
}
line := bytes.TrimSpace(w.line[:i])
w.line = w.line[i+1:]
if rest, ok := bytes.CutPrefix(line, []byte("data:")); ok && bytes.Contains(rest, []byte("usage")) {
if u, ok := parseUsage(bytes.TrimSpace(rest)); ok {
w.usage = u
}
}
}
if len(w.line) > maxCapture {
w.line = nil
}
}
// parseUsage reads chat/embeddings ("usage") and responses ("response.usage") shapes.
func parseUsage(data []byte) (Usage, bool) {
u := gjson.GetBytes(data, "usage")
if !u.IsObject() {
u = gjson.GetBytes(data, "response.usage")
}
if !u.IsObject() {
return Usage{}, false
}
first := func(keys ...string) int64 {
for _, k := range keys {
if v := u.Get(k); v.Exists() {
return v.Int()
}
}
return 0
}
out := Usage{
PromptTokens: first("prompt_tokens", "input_tokens"),
CompletionTokens: first("completion_tokens", "output_tokens"),
TotalTokens: first("total_tokens"),
}
if out.TotalTokens == 0 {
out.TotalTokens = out.PromptTokens + out.CompletionTokens
}
return out, true
}