Files
2026-09-11 06:14:38 +02:00

215 lines
5.8 KiB
Go

package cost
import (
"bytes"
"encoding/json"
"math"
"strings"
"github.com/example/ollama-fair-gateway/internal/config"
)
type Estimate struct {
Model string
InputTokens int64
OutputTokens int64
Credits float64
}
type Usage struct {
PromptTokens int64 `json:"prompt_tokens"`
CompletionTokens int64 `json:"completion_tokens"`
CachedPromptTokens int64 `json:"cached_prompt_tokens,omitempty"`
PromptEvalNS int64 `json:"prompt_eval_ns,omitempty"`
EvalNS int64 `json:"eval_ns,omitempty"`
LoadNS int64 `json:"load_ns,omitempty"`
TotalNS int64 `json:"total_ns,omitempty"`
Approximate bool `json:"approximate,omitempty"`
}
type Estimator struct{ cfg config.CostConfig }
func New(cfg config.CostConfig) *Estimator { return &Estimator{cfg: cfg} }
func (e *Estimator) Rate(model string) config.ModelRate {
if r, ok := e.cfg.Models[model]; ok {
return fillRate(r, e.cfg.Default)
}
// permit convenient prefix entries such as "qwen3:*"
best := ""
var out config.ModelRate
for p, r := range e.cfg.Models {
if strings.HasSuffix(p, "*") && strings.HasPrefix(model, strings.TrimSuffix(p, "*")) && len(p) > len(best) {
best = p
out = r
}
}
if best != "" {
return fillRate(out, e.cfg.Default)
}
return e.cfg.Default
}
func fillRate(r, d config.ModelRate) config.ModelRate {
if r.InputCreditsPer1K <= 0 {
r.InputCreditsPer1K = d.InputCreditsPer1K
}
if r.OutputCreditsPer1K <= 0 {
r.OutputCreditsPer1K = d.OutputCreditsPer1K
}
if r.CachedInputFactor <= 0 {
r.CachedInputFactor = d.CachedInputFactor
if r.CachedInputFactor <= 0 {
r.CachedInputFactor = 1
}
}
if r.ComputeCreditsPerSecond < 0 {
r.ComputeCreditsPerSecond = 0
}
return r
}
type wireReq struct {
Model string `json:"model"`
Prompt string `json:"prompt"`
Suffix string `json:"suffix"`
System json.RawMessage `json:"system"`
Instructions json.RawMessage `json:"instructions"`
Input json.RawMessage `json:"input"`
Tools json.RawMessage `json:"tools"`
Messages []struct {
Content json.RawMessage `json:"content"`
} `json:"messages"`
MaxTokens int `json:"max_tokens"`
MaxCompletionTokens int `json:"max_completion_tokens"`
MaxOutputTokens int `json:"max_output_tokens"`
Options struct {
NumPredict int `json:"num_predict"`
NumCtx int `json:"num_ctx"`
} `json:"options"`
}
func (e *Estimator) Estimate(path string, body []byte) Estimate {
var r wireReq
_ = json.Unmarshal(body, &r)
textBytes := len(r.Prompt) + len(r.Suffix) + textualBytes(r.System)
textBytes += textualBytes(r.Instructions)
textBytes += textualBytes(r.Input)
// Tool schemas are injected into the model context. Counting their raw JSON
// bytes gives a conservative pre-run estimate without attempting to tokenize
// provider-specific JSON schema dialects.
if raw := bytes.TrimSpace(r.Tools); len(raw) > 2 && !bytes.Equal(raw, []byte("[]")) && !bytes.Equal(raw, []byte("null")) {
textBytes += len(raw)
}
for _, m := range r.Messages {
textBytes += textualBytes(m.Content)
}
in := roughTokens(textBytes)
if in < 1 { // bounded fallback for unusual request shapes
n := len(body)
if n > 64<<10 {
n = 64 << 10
}
in = roughTokens(n)
}
// Use the largest explicitly requested output budget. Different compatible
// APIs use different field names, and choosing the largest is conservative
// when a client sends more than one of them.
out := maxInt(r.MaxCompletionTokens, r.MaxOutputTokens, r.MaxTokens, r.Options.NumPredict)
if out <= 0 {
out = e.cfg.DefaultMaxOutputTokens
}
if strings.Contains(path, "embeddings") || strings.Contains(path, "/embed") {
out = 0
}
rate := e.Rate(r.Model)
credits := creditsFor(rate, in, 0, int64(out), 0)
if rate.ComputeCreditsPerSecond > 0 {
var sec float64
if rate.ExpectedPromptTokensPerSecond > 0 {
sec += float64(in) / rate.ExpectedPromptTokensPerSecond
}
if out > 0 && rate.ExpectedOutputTokensPerSecond > 0 {
sec += float64(out) / rate.ExpectedOutputTokensPerSecond
}
credits += sec * rate.ComputeCreditsPerSecond
}
if credits <= 0 {
credits = .001
}
return Estimate{Model: r.Model, InputTokens: in, OutputTokens: int64(out), Credits: credits}
}
func (e *Estimator) Actual(model string, u Usage) float64 {
r := e.Rate(model)
sec := float64(u.PromptEvalNS+u.EvalNS) / 1e9
return creditsFor(r, u.PromptTokens, u.CachedPromptTokens, u.CompletionTokens, sec)
}
func creditsFor(r config.ModelRate, in, cached, out int64, sec float64) float64 {
if cached < 0 {
cached = 0
}
if cached > in {
cached = in
}
uncached := in - cached
input := (float64(uncached) + float64(cached)*r.CachedInputFactor) / 1000 * r.InputCreditsPer1K
return input + float64(out)/1000*r.OutputCreditsPer1K + sec*r.ComputeCreditsPerSecond
}
func maxInt(values ...int) int {
m := 0
for _, v := range values {
if v > m {
m = v
}
}
return m
}
func roughTokens(n int) int64 {
if n <= 0 {
return 0
}
return int64(math.Ceil(float64(n) / 4))
}
func textualBytes(raw json.RawMessage) int {
raw = bytes.TrimSpace(raw)
if len(raw) == 0 || bytes.Equal(raw, []byte("null")) {
return 0
}
var s string
if raw[0] == '"' && json.Unmarshal(raw, &s) == nil {
return len(s)
}
var v any
dec := json.NewDecoder(bytes.NewReader(raw))
if dec.Decode(&v) != nil {
return 0
}
return walkText(v)
}
func walkText(v any) int {
switch x := v.(type) {
case string:
return len(x)
case []any:
n := 0
for _, z := range x {
n += walkText(z)
}
return n
case map[string]any:
// Follow only known textual fields. URLs/base64 image payloads are
// deliberately excluded from the pre-run token estimate.
n := 0
if v, ok := x["text"]; ok {
n += walkText(v)
}
if v, ok := x["content"]; ok {
n += walkText(v)
}
return n
}
return 0
}