Skip to content

Commit 2014ff9

Browse files
authored
Update metrics.py
1 parent 3eaadcf commit 2014ff9

1 file changed

Lines changed: 174 additions & 116 deletions

File tree

‎app/infra/metrics.py‎

Lines changed: 174 additions & 116 deletions
Original file line numberDiff line numberDiff line change
@@ -1,122 +1,180 @@
1-
"""Usage business logic: aggregate the token ledger for a user."""
1+
"""In-process metrics registry + tracing hooks.
22
3-
from __future__ import annotations
4-
5-
from dataclasses import dataclass, field
6-
7-
from sqlalchemy import func, select
8-
from sqlalchemy.orm import Session
9-
10-
from ..infra.config import get_settings
11-
from ..models import Conversation, Run, UsageEvent, User
12-
from .errors import PaymentRequired
13-
14-
15-
@dataclass
16-
class UsageTotals:
17-
"""A user's billed totals, plus a per-conversation token breakdown."""
18-
19-
input_tokens: int = 0
20-
output_tokens: int = 0
21-
rounds: int = 0
22-
runs: int = 0
23-
by_conversation: dict[str, int] = field(default_factory=dict)
3+
Dependency-free by design: a small counter/histogram registry that
4+
renders Prometheus text exposition at ``/metrics``, plus a ``span``
5+
tracing hook that is a no-op until a tracer is registered (so wiring
6+
OpenTelemetry later is a one-liner, and nothing is required today).
247
8+
Not a full metrics client -- it is the minimum that makes runs, tokens,
9+
and HTTP observable by a scraper, without pulling in a backend. Labels
10+
are supported as sorted key/value tuples so a metric can be sliced
11+
(e.g. http_requests_total by method+status).
12+
"""
2513

26-
def summary(db: Session, user: User) -> UsageTotals:
27-
"""Totals over the user's usage events, grouped per conversation.
14+
from __future__ import annotations
2815

29-
``by_conversation`` carries combined input+output tokens, which is
30-
what a bill is drawn on; the run count comes from the Run table so
31-
runs that recorded no usage (failures) are still counted.
16+
import threading
17+
import time
18+
from collections.abc import Iterator
19+
from contextlib import contextmanager
20+
from typing import Any, Protocol
21+
22+
_Labels = tuple[tuple[str, str], ...]
23+
24+
25+
def _labelset(labels: dict[str, str] | None) -> _Labels:
26+
if not labels:
27+
return ()
28+
return tuple(sorted(labels.items()))
29+
30+
31+
# Prometheus default histogram buckets (seconds), fine for request/run latency.
32+
_DEFAULT_BUCKETS = (0.005, 0.025, 0.1, 0.5, 1.0, 5.0, 30.0, 120.0)
33+
34+
35+
class MetricsRegistry:
36+
"""Thread-safe counters and histograms with optional labels."""
37+
38+
def __init__(self) -> None:
39+
self._lock = threading.Lock()
40+
self._counters: dict[tuple[str, _Labels], float] = {}
41+
self._hist_sum: dict[tuple[str, _Labels], float] = {}
42+
self._hist_count: dict[tuple[str, _Labels], int] = {}
43+
self._hist_buckets: dict[tuple[str, _Labels], list[int]] = {}
44+
self._help: dict[str, str] = {}
45+
46+
def counter(
47+
self, name: str, value: float = 1.0, labels: dict[str, str] | None = None, help: str = ""
48+
) -> None:
49+
key = (name, _labelset(labels))
50+
with self._lock:
51+
self._counters[key] = self._counters.get(key, 0.0) + value
52+
if help and name not in self._help:
53+
self._help[name] = help
54+
55+
def observe(
56+
self, name: str, value: float, labels: dict[str, str] | None = None, help: str = ""
57+
) -> None:
58+
key = (name, _labelset(labels))
59+
with self._lock:
60+
self._hist_sum[key] = self._hist_sum.get(key, 0.0) + value
61+
self._hist_count[key] = self._hist_count.get(key, 0) + 1
62+
buckets = self._hist_buckets.get(key)
63+
if buckets is None:
64+
buckets = [0] * len(_DEFAULT_BUCKETS)
65+
self._hist_buckets[key] = buckets
66+
for i, edge in enumerate(_DEFAULT_BUCKETS):
67+
if value <= edge:
68+
buckets[i] += 1
69+
if help and name not in self._help:
70+
self._help[name] = help
71+
72+
def reset(self) -> None:
73+
with self._lock:
74+
self._counters.clear()
75+
self._hist_sum.clear()
76+
self._hist_count.clear()
77+
self._hist_buckets.clear()
78+
79+
def render(self) -> str:
80+
"""Prometheus text exposition format."""
81+
lines: list[str] = []
82+
with self._lock:
83+
for name in sorted({n for (n, _) in self._counters}):
84+
if name in self._help:
85+
lines.append(f"# HELP {name} {self._help[name]}")
86+
lines.append(f"# TYPE {name} counter")
87+
for (n, labels), val in self._counters.items():
88+
if n != name:
89+
continue
90+
lines.append(f"{name}{_fmt_labels(labels)} {val}")
91+
for name in sorted({n for (n, _) in self._hist_count}):
92+
if name in self._help:
93+
lines.append(f"# HELP {name} {self._help[name]}")
94+
lines.append(f"# TYPE {name} histogram")
95+
for (n, labels), buckets in self._hist_buckets.items():
96+
if n != name:
97+
continue
98+
cumulative = 0
99+
for i, edge in enumerate(_DEFAULT_BUCKETS):
100+
cumulative += buckets[i]
101+
le = _fmt_labels(labels, extra=("le", _fmt_float(edge)))
102+
lines.append(f"{name}_bucket{le} {cumulative}")
103+
inf = _fmt_labels(labels, extra=("le", "+Inf"))
104+
lines.append(f"{name}_bucket{inf} {self._hist_count[(n, labels)]}")
105+
lines.append(f"{name}_sum{_fmt_labels(labels)} {self._hist_sum[(n, labels)]}")
106+
lines.append(
107+
f"{name}_count{_fmt_labels(labels)} {self._hist_count[(n, labels)]}"
108+
)
109+
return "\n".join(lines) + "\n"
110+
111+
112+
def _fmt_float(v: float) -> str:
113+
return repr(v)
114+
115+
116+
def _fmt_labels(labels: _Labels, extra: tuple[str, str] | None = None) -> str:
117+
items = list(labels)
118+
if extra is not None:
119+
items = items + [extra]
120+
if not items:
121+
return ""
122+
inner = ",".join(f'{k}="{v}"' for k, v in items)
123+
return "{" + inner + "}"
124+
125+
126+
_registry = MetricsRegistry()
127+
128+
129+
def get_registry() -> MetricsRegistry:
130+
return _registry
131+
132+
133+
# Bounded set of HTTP methods for the method label -- an arbitrary or
134+
# garbage verb must not create a new permanent metric series (unbounded
135+
# label cardinality is a memory-growth DoS). Anything else is "other".
136+
_KNOWN_METHODS = frozenset({"GET", "POST", "PUT", "PATCH", "DELETE", "HEAD", "OPTIONS"})
137+
138+
139+
def normalize_method(method: str) -> str:
140+
"""Map an HTTP method to a bounded label value."""
141+
return method.upper() if method.upper() in _KNOWN_METHODS else "other"
142+
143+
144+
# -- tracing hooks --------------------------------------------------------
145+
146+
147+
class Tracer(Protocol):
148+
@contextmanager
149+
def span(self, name: str, attributes: dict[str, Any] | None = None) -> Iterator[None]: ...
150+
151+
152+
_tracer: Tracer | None = None
153+
154+
155+
def set_tracer(tracer: Tracer | None) -> None:
156+
"""Register a tracer (e.g. an OpenTelemetry adapter). ``None``
157+
restores the no-op behavior."""
158+
global _tracer
159+
_tracer = tracer
160+
161+
162+
@contextmanager
163+
def span(name: str, attributes: dict[str, Any] | None = None) -> Iterator[None]:
164+
"""Trace a block of work.
165+
166+
A no-op (zero overhead beyond the context manager) unless a tracer
167+
is registered via ``set_tracer`` -- so run/request spans are already
168+
marked in the code and become real spans the moment a backend is
169+
wired, with no call-site changes. Also records a duration histogram
170+
so timing is visible even without a tracer.
32171
"""
33-
totals = UsageTotals()
34-
rows = db.execute(
35-
select(
36-
UsageEvent.conversation_id,
37-
func.sum(UsageEvent.input_tokens),
38-
func.sum(UsageEvent.output_tokens),
39-
func.sum(UsageEvent.rounds),
40-
)
41-
.where(UsageEvent.user_id == user.id)
42-
.group_by(UsageEvent.conversation_id)
43-
).all()
44-
for conversation_id, inp, out, rounds in rows:
45-
inp_i, out_i, rounds_i = int(inp or 0), int(out or 0), int(rounds or 0)
46-
totals.input_tokens += inp_i
47-
totals.output_tokens += out_i
48-
totals.rounds += rounds_i
49-
totals.by_conversation[conversation_id] = inp_i + out_i
50-
totals.runs = int(
51-
db.query(func.count(Run.id))
52-
.join(Conversation, Run.conversation_id == Conversation.id)
53-
.filter(Conversation.user_id == user.id)
54-
.scalar()
55-
or 0
172+
start = time.perf_counter()
173+
if _tracer is not None:
174+
with _tracer.span(name, attributes):
175+
yield
176+
else:
177+
yield
178+
_registry.observe(
179+
"paw_span_duration_seconds", time.perf_counter() - start, labels={"span": name}
56180
)
57-
return totals
58-
59-
60-
def consumed_tokens(db: Session, user: User) -> int:
61-
"""Total tokens (input + output) the user has spent, from the ledger."""
62-
row = db.execute(
63-
select(
64-
func.coalesce(func.sum(UsageEvent.input_tokens), 0)
65-
+ func.coalesce(func.sum(UsageEvent.output_tokens), 0)
66-
).where(UsageEvent.user_id == user.id)
67-
).scalar()
68-
return int(row or 0)
69-
70-
71-
def token_budget(user: User) -> int:
72-
"""The user's total token allowance.
73-
74-
Free-tier allowance plus the token value of their purchased points
75-
(both buckets). A points top-up therefore raises the ceiling.
76-
"""
77-
s = get_settings()
78-
points = (user.plan_points or 0) + (user.pack_points or 0)
79-
return s.budget_free_tokens + points * s.budget_tokens_per_point
80-
81-
82-
@dataclass(frozen=True)
83-
class BudgetStatus:
84-
consumed: int
85-
budget: int
86-
87-
@property
88-
def remaining(self) -> int:
89-
return max(0, self.budget - self.consumed)
90-
91-
@property
92-
def exhausted(self) -> bool:
93-
return self.consumed >= self.budget
94-
95-
96-
def budget_status(db: Session, user: User) -> BudgetStatus:
97-
return BudgetStatus(consumed=consumed_tokens(db, user), budget=token_budget(user))
98-
99-
100-
def enforce_budget(db: Session, user: User) -> None:
101-
"""Kill-switch: refuse a new run when the user is out of budget.
102-
103-
A no-op unless ``budget_enforce`` is on, so dev/internal deployments
104-
are unaffected. Enforced *before* a run starts against a snapshot of
105-
the ledger.
106-
107-
Bound, stated honestly: the ledger only records after a run finishes,
108-
and the per-conversation running-run guard admits one concurrent run
109-
*per conversation*. So a user with many conversations can start one
110-
run per conversation against the same pre-spend snapshot -- overspend
111-
is bounded by the number of the user's conversations, not by one run
112-
globally. Tightening that to a hard global cap would require
113-
reserving tokens up front (a debit-on-start ledger), which is a
114-
deliberate future step, not implemented here.
115-
"""
116-
if not get_settings().budget_enforce:
117-
return
118-
status = budget_status(db, user)
119-
if status.exhausted:
120-
raise PaymentRequired(
121-
f"token budget exhausted ({status.consumed}/{status.budget}); add credits to continue"
122-
)

0 commit comments

Comments
 (0)