refactor(gateway): deduplicate stats CSV export paths

- extract csvHeaders() (Content-Type + Content-Disposition) shared by both
  export branches
- extract exportKey(): the identical admin/user key-filter logic existed
  twice (JSON path + keys-csv); now all three call sites share one function
- inline the nine single-use intermediate variables in the keys-csv loop
This commit is contained in:
dev
2026-08-24 22:41:48 +08:00
parent 98847adfd6
commit f3d5ba6cea

View File

@ -128,6 +128,23 @@ func (g *Gateway) handleSourcesAPI(w http.ResponseWriter, r *http.Request) {
} }
} }
// csvHeaders stamps the shared download headers for CSV exports.
func csvHeaders(w http.ResponseWriter, filename string) {
w.Header().Set("Content-Type", "text/csv; charset=utf-8")
w.Header().Set("Content-Disposition", "attachment; filename="+filename)
}
// exportKey resolves which masked key id an export covers: admins may pass
// any key filter, user keys are always scoped to themselves (records and
// aggregates are keyed by the masked keyID form).
func exportKey(r *http.Request) string {
k := r.URL.Query().Get("key")
if reqRole(r.Context()) != "admin" {
k = keyID(reqKey(r.Context()))
}
return k
}
// handleStatsAPI returns per-key / per-model / per-source usage aggregates and // handleStatsAPI returns per-key / per-model / per-source usage aggregates and
// the recent request audit trail. // the recent request audit trail.
func (g *Gateway) handleStatsAPI(w http.ResponseWriter, r *http.Request) { func (g *Gateway) handleStatsAPI(w http.ResponseWriter, r *http.Request) {
@ -141,20 +158,14 @@ func (g *Gateway) handleStatsAPI(w http.ResponseWriter, r *http.Request) {
limit = n limit = n
} }
} }
key := r.URL.Query().Get("key") key := exportKey(r)
if reqRole(r.Context()) != "admin" {
// user keys may only see their own usage - records and aggregates are
// keyed by the masked id (keyID), so filter on that masked form.
key = keyID(reqKey(r.Context()))
}
if r.URL.Query().Get("export") == "csv" { if r.URL.Query().Get("export") == "csv" {
from, _ := strconv.ParseInt(r.URL.Query().Get("from"), 10, 64) from, _ := strconv.ParseInt(r.URL.Query().Get("from"), 10, 64)
to, _ := strconv.ParseInt(r.URL.Query().Get("to"), 10, 64) to, _ := strconv.ParseInt(r.URL.Query().Get("to"), 10, 64)
if to == 0 { if to == 0 {
to = time.Now().UnixMilli() to = time.Now().UnixMilli()
} }
w.Header().Set("Content-Type", "text/csv; charset=utf-8") csvHeaders(w, "llmsproxy-requests.csv")
w.Header().Set("Content-Disposition", "attachment; filename=llmsproxy-requests.csv")
cw := csv.NewWriter(w) cw := csv.NewWriter(w)
names := map[string]string{} names := map[string]string{}
for _, k := range g.core.ListKeys() { for _, k := range g.core.ListKeys() {
@ -181,31 +192,24 @@ func (g *Gateway) handleStatsAPI(w http.ResponseWriter, r *http.Request) {
return return
} }
if r.URL.Query().Get("export") == "keys-csv" { if r.URL.Query().Get("export") == "keys-csv" {
w.Header().Set("Content-Type", "text/csv; charset=utf-8") csvHeaders(w, "llmsproxy-keys.csv")
w.Header().Set("Content-Disposition", "attachment; filename=llmsproxy-keys.csv")
cw := csv.NewWriter(w) cw := csv.NewWriter(w)
_ = cw.Write([]string{"key", "key_name", "role", "models", "total_requests", "success_requests", "failed_requests", "prompt_tokens", "completion_tokens", "total_tokens", "avg_latency_ms", "max_latency_ms", "created_at"}) _ = cw.Write([]string{"key", "key_name", "role", "models", "total_requests", "success_requests", "failed_requests", "prompt_tokens", "completion_tokens", "total_tokens", "avg_latency_ms", "max_latency_ms", "created_at"})
// Use the same key filtering as the JSON API keyFilter := exportKey(r)
exportKey := r.URL.Query().Get("key")
if reqRole(r.Context()) != "admin" {
exportKey = keyID(reqKey(r.Context()))
}
// by_key rows are keyed by the masked id (keyID); build a masked-id -> // by_key rows are keyed by the masked id (keyID); build a masked-id ->
// record map so names/roles/models resolve for the export. // record map so names/roles/models resolve for the export.
byMasked := map[string]config.GWKey{} byMasked := map[string]config.GWKey{}
for _, k := range g.core.ListKeys() { for _, k := range g.core.ListKeys() {
byMasked[keyID(k.Key)] = k byMasked[keyID(k.Key)] = k
} }
snap := g.stats.Snapshot(0, exportKey) snap := g.stats.Snapshot(0, keyFilter)
byKeyRaw, ok := snap["by_key"].([]StatsRow) byKeyRaw, ok := snap["by_key"].([]StatsRow)
if !ok { if !ok {
byKeyRaw = []StatsRow{} byKeyRaw = []StatsRow{}
} }
for _, row := range byKeyRaw { for _, row := range byKeyRaw {
keyInfo, found := byMasked[row.Name] keyInfo, found := byMasked[row.Name]
name := "" name, role, models := "", "", ""
role := ""
models := ""
if found { if found {
name = keyInfo.Name name = keyInfo.Name
role = keyInfo.Role role = keyInfo.Role
@ -215,17 +219,9 @@ func (g *Gateway) handleStatsAPI(w http.ResponseWriter, r *http.Request) {
} }
models = strings.Join(modelNames, ",") models = strings.Join(modelNames, ",")
} }
reqs := row.Stat.Reqs
success := row.Stat.OK
errCount := row.Stat.Err
prompt := row.Stat.Prompt
compl := row.Stat.Compl
tokens := row.Stat.Prompt + row.Stat.Compl
latSum := row.Stat.LatSum
latMax := row.Stat.LatMax
avgLat := int64(0) avgLat := int64(0)
if reqs > 0 { if row.Stat.Reqs > 0 {
avgLat = latSum / reqs avgLat = row.Stat.LatSum / row.Stat.Reqs
} }
createdAt := "" createdAt := ""
if found && keyInfo.CreatedAt > 0 { if found && keyInfo.CreatedAt > 0 {
@ -233,14 +229,14 @@ func (g *Gateway) handleStatsAPI(w http.ResponseWriter, r *http.Request) {
} }
_ = cw.Write([]string{ _ = cw.Write([]string{
row.Name, name, role, models, row.Name, name, role, models,
strconv.FormatInt(reqs, 10), strconv.FormatInt(row.Stat.Reqs, 10),
strconv.FormatInt(success, 10), strconv.FormatInt(row.Stat.OK, 10),
strconv.FormatInt(errCount, 10), strconv.FormatInt(row.Stat.Err, 10),
strconv.FormatInt(prompt, 10), strconv.FormatInt(row.Stat.Prompt, 10),
strconv.FormatInt(compl, 10), strconv.FormatInt(row.Stat.Compl, 10),
strconv.FormatInt(tokens, 10), strconv.FormatInt(row.Stat.Prompt+row.Stat.Compl, 10),
strconv.FormatInt(avgLat, 10), strconv.FormatInt(avgLat, 10),
strconv.FormatInt(latMax, 10), strconv.FormatInt(row.Stat.LatMax, 10),
createdAt, createdAt,
}) })
} }