mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 03:12:55 +00:00
feat(quotas): quota_policies table and a per-call cost on token_usage
Migration 0033 adds quota_policies (instance default, team per-member allowance, user override; a token budget and a USD budget per row) and token_usage.cost. The column is added IF NOT EXISTS so a database that already carries it upgrades cleanly. Every usage row now records the call's USD cost from the model catalog; bring-your-own models are recorded at $0.
This commit is contained in:
1 parent
b77561288d
commit
43fad2a865
9 files changed
+415
-3
No files matched your search
@@ -0,0 +1,102 @@
|
||||
"""0033 quotas — admin-set usage limits and a per-call cost.
|
||||
|
||||
``quota_policies`` holds the limits an instance admin sets at three layers:
|
||||
the instance default (``subject_id`` NULL), a team's per-member allowance
|
||||
(``subject_id`` = ``teams.id``) and a single user's override (``subject_id`` =
|
||||
the auth ``sub``). Each row carries a token budget and a USD budget; per
|
||||
budget a row either sets a limit (0 blocks), marks it unlimited, or leaves
|
||||
both empty to defer to the next layer. ``bucket`` narrows a row to chat without
|
||||
an agent (``direct``) or traffic through an agent (``agent``); ``all`` covers both.
|
||||
|
||||
``subject_id`` is polymorphic, so there is no FK: an AFTER DELETE trigger on
|
||||
``teams`` scrubs a deleted team's rows, and user rows follow the ``user_roles``
|
||||
convention of never blocking user deletion.
|
||||
|
||||
``token_usage.cost`` is the USD cost of the call at write time (see
|
||||
``docsgpt/pricing.py``); 0 for unpriced and bring-your-own models.
|
||||
Idempotent both ways.
|
||||
|
||||
Revision ID: 0033_quotas
|
||||
Revises: 0032_personal_access_tokens
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "0033_quotas"
|
||||
down_revision: Union[str, None] = "0032_personal_access_tokens"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.execute(
|
||||
"ALTER TABLE token_usage ADD COLUMN IF NOT EXISTS cost NUMERIC(12,8) NOT NULL DEFAULT 0;"
|
||||
)
|
||||
op.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS quota_policies (
|
||||
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||
scope TEXT NOT NULL CHECK (scope IN ('instance', 'team', 'user')),
|
||||
subject_id TEXT,
|
||||
bucket TEXT NOT NULL DEFAULT 'all'
|
||||
CHECK (bucket IN ('all', 'direct', 'agent')),
|
||||
token_limit BIGINT CHECK (token_limit >= 0),
|
||||
token_unlimited BOOLEAN NOT NULL DEFAULT false,
|
||||
cost_limit_usd NUMERIC(12,4) CHECK (cost_limit_usd >= 0),
|
||||
cost_unlimited BOOLEAN NOT NULL DEFAULT false,
|
||||
enabled BOOLEAN NOT NULL DEFAULT true,
|
||||
note TEXT,
|
||||
created_by TEXT,
|
||||
updated_by TEXT,
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||
updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
|
||||
CONSTRAINT quota_policies_subject_chk
|
||||
CHECK ((scope = 'instance') = (subject_id IS NULL)),
|
||||
CONSTRAINT quota_policies_token_chk
|
||||
CHECK (NOT (token_unlimited AND token_limit IS NOT NULL)),
|
||||
CONSTRAINT quota_policies_cost_chk
|
||||
CHECK (NOT (cost_unlimited AND cost_limit_usd IS NOT NULL))
|
||||
);
|
||||
"""
|
||||
)
|
||||
# One row per (layer subject, bucket); the instance row's NULL subject folds to ''.
|
||||
op.execute(
|
||||
"CREATE UNIQUE INDEX IF NOT EXISTS quota_policies_subject_uidx "
|
||||
"ON quota_policies (scope, COALESCE(subject_id, ''), bucket);"
|
||||
)
|
||||
op.execute("DROP TRIGGER IF EXISTS quota_policies_set_updated_at ON quota_policies;")
|
||||
op.execute(
|
||||
"""
|
||||
CREATE TRIGGER quota_policies_set_updated_at
|
||||
BEFORE UPDATE ON quota_policies
|
||||
FOR EACH ROW EXECUTE FUNCTION set_updated_at();
|
||||
"""
|
||||
)
|
||||
op.execute(
|
||||
"""
|
||||
CREATE OR REPLACE FUNCTION cleanup_team_quota_policies() RETURNS trigger AS $$
|
||||
BEGIN
|
||||
DELETE FROM quota_policies WHERE scope = 'team' AND subject_id = OLD.id::text;
|
||||
RETURN OLD;
|
||||
END;
|
||||
$$ LANGUAGE plpgsql;
|
||||
"""
|
||||
)
|
||||
op.execute("DROP TRIGGER IF EXISTS teams_cleanup_quota_policies ON teams;")
|
||||
op.execute(
|
||||
"""
|
||||
CREATE TRIGGER teams_cleanup_quota_policies
|
||||
AFTER DELETE ON teams
|
||||
FOR EACH ROW EXECUTE FUNCTION cleanup_team_quota_policies();
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.execute("DROP TRIGGER IF EXISTS teams_cleanup_quota_policies ON teams;")
|
||||
op.execute("DROP FUNCTION IF EXISTS cleanup_team_quota_policies();")
|
||||
op.execute("DROP TABLE IF EXISTS quota_policies;")
|
||||
op.execute("ALTER TABLE token_usage DROP COLUMN IF EXISTS cost;")
|
||||
@@ -52,6 +52,7 @@ class LLMCreator:
|
||||
base_url = None
|
||||
upstream_model_id = model_id
|
||||
capabilities = None
|
||||
model = None
|
||||
if model_id:
|
||||
user_id = model_user_id
|
||||
if user_id is None:
|
||||
@@ -127,4 +128,6 @@ class LLMCreator:
|
||||
# llm.model_id is the upstream name (BYOM resolves it above); stamp
|
||||
# the canonical id (UUID for BYOM) separately for token_usage.
|
||||
llm._canonical_model_id = model_id
|
||||
# Calls to a user's own model are recorded at $0 (see ``docsgpt/usage.py``).
|
||||
llm._is_byom = model is not None and getattr(model, "source", "builtin") == "user"
|
||||
return llm
|
||||
@@ -28,6 +28,7 @@ from sqlalchemy import (
|
||||
Index,
|
||||
Integer,
|
||||
MetaData,
|
||||
Numeric,
|
||||
PrimaryKeyConstraint,
|
||||
UniqueConstraint,
|
||||
Table,
|
||||
@@ -241,6 +242,9 @@ token_usage_table = Table(
|
||||
# cache activity, so hit-rate queries stay honest across providers.
|
||||
Column("cached_tokens", Integer),
|
||||
Column("cache_write_tokens", Integer),
|
||||
# Added in ``0033_quotas``. USD cost of the call at write time; 0 for
|
||||
# unpriced and bring-your-own models.
|
||||
Column("cost", Numeric(12, 8), nullable=False, server_default="0"),
|
||||
)
|
||||
|
||||
user_logs_table = Table(
|
||||
@@ -1121,3 +1125,42 @@ Index(
|
||||
personal_access_tokens_table.c.user_id,
|
||||
personal_access_tokens_table.c.created_at.desc(),
|
||||
)
|
||||
|
||||
# --- Usage quotas (migration 0033) ------------------------------------------
|
||||
# Admin-set limits at three layers: instance (``subject_id`` NULL), team
|
||||
# per-member allowance (``teams.id``) and user override (auth ``sub``). Per
|
||||
# budget a row sets a limit, marks it unlimited, or defers to the next layer.
|
||||
|
||||
quota_policies_table = Table(
|
||||
"quota_policies",
|
||||
metadata,
|
||||
Column("id", UUID(as_uuid=True), primary_key=True, server_default=func.gen_random_uuid()),
|
||||
Column("scope", Text, nullable=False),
|
||||
Column("subject_id", Text),
|
||||
Column("bucket", Text, nullable=False, server_default="all"),
|
||||
Column("token_limit", BigInteger),
|
||||
Column("token_unlimited", Boolean, nullable=False, server_default="false"),
|
||||
Column("cost_limit_usd", Numeric(12, 4)),
|
||||
Column("cost_unlimited", Boolean, nullable=False, server_default="false"),
|
||||
Column("enabled", Boolean, nullable=False, server_default="true"),
|
||||
Column("note", Text),
|
||||
Column("created_by", Text),
|
||||
Column("updated_by", Text),
|
||||
Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
Column("updated_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
CheckConstraint("scope IN ('instance', 'team', 'user')", name="quota_policies_scope_check"),
|
||||
CheckConstraint("bucket IN ('all', 'direct', 'agent')", name="quota_policies_bucket_check"),
|
||||
CheckConstraint("token_limit >= 0", name="quota_policies_token_limit_check"),
|
||||
CheckConstraint("cost_limit_usd >= 0", name="quota_policies_cost_limit_check"),
|
||||
CheckConstraint("(scope = 'instance') = (subject_id IS NULL)", name="quota_policies_subject_chk"),
|
||||
CheckConstraint("NOT (token_unlimited AND token_limit IS NOT NULL)", name="quota_policies_token_chk"),
|
||||
CheckConstraint("NOT (cost_unlimited AND cost_limit_usd IS NOT NULL)", name="quota_policies_cost_chk"),
|
||||
)
|
||||
|
||||
Index(
|
||||
"quota_policies_subject_uidx",
|
||||
quota_policies_table.c.scope,
|
||||
func.coalesce(quota_policies_table.c.subject_id, ""),
|
||||
quota_policies_table.c.bucket,
|
||||
unique=True,
|
||||
)
|
||||
@@ -37,6 +37,7 @@ class TokenUsageRepository:
|
||||
timestamp: Optional[datetime] = None,
|
||||
cached_tokens: Optional[int] = None,
|
||||
cache_write_tokens: Optional[int] = None,
|
||||
cost: float = 0.0,
|
||||
) -> None:
|
||||
# Attribution guard: the ``token_usage_attribution_chk`` CHECK
|
||||
# constraint requires at least one of ``user_id`` / ``api_key``
|
||||
@@ -62,14 +63,14 @@ class TokenUsageRepository:
|
||||
INSERT INTO token_usage (
|
||||
user_id, api_key, agent_id,
|
||||
prompt_tokens, generated_tokens,
|
||||
cached_tokens, cache_write_tokens,
|
||||
cached_tokens, cache_write_tokens, cost,
|
||||
source, request_id, model_id, timestamp
|
||||
)
|
||||
VALUES (
|
||||
:user_id, :api_key,
|
||||
CAST(:agent_id AS uuid),
|
||||
:prompt_tokens, :generated_tokens,
|
||||
:cached_tokens, :cache_write_tokens,
|
||||
:cached_tokens, :cache_write_tokens, :cost,
|
||||
:source, :request_id, :model_id, COALESCE(:timestamp, now())
|
||||
)
|
||||
"""
|
||||
@@ -82,6 +83,7 @@ class TokenUsageRepository:
|
||||
"generated_tokens": generated_tokens,
|
||||
"cached_tokens": cached_tokens,
|
||||
"cache_write_tokens": cache_write_tokens,
|
||||
"cost": cost,
|
||||
"source": source,
|
||||
"request_id": request_id,
|
||||
"model_id": model_id,
|
||||
|
||||
+24
-1
@@ -2,6 +2,7 @@ import logging
|
||||
import time
|
||||
from typing import Any, Dict
|
||||
|
||||
from docsgpt.pricing import compute_cost_usd
|
||||
from docsgpt.storage.db.repositories.token_usage import TokenUsageRepository
|
||||
from docsgpt.storage.db.session import db_session
|
||||
from docsgpt.utils import num_tokens_from_object_or_list, num_tokens_from_string
|
||||
@@ -119,6 +120,12 @@ def _persist_call_usage(llm, call_usage):
|
||||
},
|
||||
)
|
||||
return
|
||||
model_id = getattr(llm, "_canonical_model_id", None)
|
||||
# Bring-your-own models run on the user's own provider key: recorded, never priced.
|
||||
if getattr(llm, "_is_byom", False):
|
||||
cost = 0.0
|
||||
else:
|
||||
cost = _call_cost_usd(model_id, call_usage)
|
||||
try:
|
||||
with db_session() as conn:
|
||||
# ``timestamp`` is omitted so Postgres ``server_default
|
||||
@@ -136,16 +143,32 @@ def _persist_call_usage(llm, call_usage):
|
||||
# "0% cache hits".
|
||||
cached_tokens=call_usage.get("cached_tokens"),
|
||||
cache_write_tokens=call_usage.get("cache_write_tokens"),
|
||||
cost=cost,
|
||||
source=(
|
||||
getattr(llm, "_token_usage_source", None) or "agent_stream"
|
||||
),
|
||||
request_id=getattr(llm, "_request_id", None),
|
||||
model_id=getattr(llm, "_canonical_model_id", None),
|
||||
model_id=model_id,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("token_usage persist failed")
|
||||
|
||||
|
||||
def _call_cost_usd(model_id, call_usage) -> float:
|
||||
"""Price one call; a pricing failure records $0 rather than dropping the row."""
|
||||
try:
|
||||
return compute_cost_usd(
|
||||
model_id,
|
||||
call_usage["prompt_tokens"],
|
||||
call_usage["generated_tokens"],
|
||||
cached_tokens=call_usage.get("cached_tokens"),
|
||||
cache_write_tokens=call_usage.get("cache_write_tokens"),
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("token_usage cost computation failed")
|
||||
return 0.0
|
||||
|
||||
|
||||
def _prefer_provider_usage(llm: Any, call_usage: Dict[str, int]) -> Dict[str, int]:
|
||||
"""Replace estimates with upstream counts when a provider reported them.
|
||||
|
||||
|
||||
@@ -1015,6 +1015,38 @@ class TestLLMCreatorPassesModelUserId:
|
||||
assert captured["model_user_id"] == "owner-alice"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_id, source, expected",
|
||||
[(None, None, False), ("catalog-model", "builtin", False), ("byom-uuid", "user", True)],
|
||||
)
|
||||
def test_byom_flag_follows_the_model_source(self, monkeypatch, model_id, source, expected):
|
||||
from types import SimpleNamespace
|
||||
|
||||
from docsgpt.llm.llm_creator import LLMCreator
|
||||
from docsgpt.llm.providers import PROVIDERS_BY_NAME
|
||||
|
||||
class _LLM:
|
||||
def __init__(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
monkeypatch.setattr(PROVIDERS_BY_NAME["openai"], "llm_class", _LLM)
|
||||
model = SimpleNamespace(
|
||||
source=source, api_key="own-key", base_url=None, upstream_model_id=None, capabilities=None
|
||||
)
|
||||
registry = SimpleNamespace(get_model=lambda _id, user_id=None: model)
|
||||
monkeypatch.setattr(
|
||||
"docsgpt.core.model_registry.ModelRegistry.get_instance", lambda: registry
|
||||
)
|
||||
|
||||
llm = LLMCreator.create_llm(
|
||||
type="openai", api_key="k", user_api_key=None,
|
||||
decoded_token={"sub": "u1"}, model_id=model_id,
|
||||
)
|
||||
|
||||
assert llm._is_byom is expected
|
||||
assert llm._canonical_model_id == model_id
|
||||
|
||||
|
||||
# Tests — responding-provider tracking (cross-provider fallback handler fix)
|
||||
|
||||
|
||||
|
||||
@@ -35,6 +35,15 @@ class TestInsert:
|
||||
)
|
||||
assert total == 30
|
||||
|
||||
def test_cost_defaults_to_zero_and_round_trips(self, pg_conn):
|
||||
repo = _repo(pg_conn)
|
||||
repo.insert(user_id="u-cost", prompt_tokens=1, generated_tokens=1)
|
||||
repo.insert(user_id="u-cost", prompt_tokens=1, generated_tokens=1, cost=0.00012345)
|
||||
costs = pg_conn.execute(
|
||||
text("SELECT cost FROM token_usage WHERE user_id = 'u-cost' ORDER BY id")
|
||||
).scalars().all()
|
||||
assert [float(c) for c in costs] == [0.0, 0.00012345]
|
||||
|
||||
|
||||
class TestReassignApiKey:
|
||||
def test_rewrites_and_preserves_rate_limit_window(self, pg_conn):
|
||||
|
||||
@@ -0,0 +1,145 @@
|
||||
"""Migration round-trip test for 0033_quotas."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
|
||||
pytestmark = pytest.mark.integration
|
||||
|
||||
|
||||
def _alembic_ini() -> Path:
|
||||
return Path(__file__).resolve().parents[3] / "docsgpt" / "alembic.ini"
|
||||
|
||||
|
||||
def _run_alembic(url: str, *args: str) -> None:
|
||||
subprocess.check_call(
|
||||
[sys.executable, "-m", "alembic", "-c", str(_alembic_ini()), *args],
|
||||
timeout=60,
|
||||
env={**os.environ, "POSTGRES_URI": url},
|
||||
)
|
||||
|
||||
|
||||
def _alembic_heads(url: str) -> list[str]:
|
||||
out = subprocess.check_output(
|
||||
[sys.executable, "-m", "alembic", "-c", str(_alembic_ini()), "heads"],
|
||||
timeout=60,
|
||||
env={**os.environ, "POSTGRES_URI": url},
|
||||
text=True,
|
||||
)
|
||||
return [line for line in out.splitlines() if line.strip()]
|
||||
|
||||
|
||||
def _alembic_version(conn) -> str:
|
||||
return conn.execute(text("SELECT version_num FROM alembic_version")).scalar()
|
||||
|
||||
|
||||
def _column_exists(conn, table: str, column: str) -> bool:
|
||||
row = conn.execute(
|
||||
text(
|
||||
"SELECT 1 FROM information_schema.columns "
|
||||
"WHERE table_name = :t AND column_name = :c AND table_schema = 'public'"
|
||||
),
|
||||
{"t": table, "c": column},
|
||||
).fetchone()
|
||||
return row is not None
|
||||
|
||||
|
||||
_0033 = "0033_quotas"
|
||||
_0032 = "0032_personal_access_tokens"
|
||||
|
||||
|
||||
def _table_exists(conn, table: str) -> bool:
|
||||
return conn.execute(text("SELECT to_regclass(:t)"), {"t": f"public.{table}"}).scalar() is not None
|
||||
|
||||
|
||||
def _insert_policy(conn, **values) -> None:
|
||||
cols = ", ".join(values)
|
||||
params = ", ".join(f":{k}" for k in values)
|
||||
conn.execute(text(f"INSERT INTO quota_policies ({cols}) VALUES ({params})"), values)
|
||||
|
||||
|
||||
class TestMigration0033RoundTrip:
|
||||
def test_single_head(self, pg_engine):
|
||||
url = pg_engine.url.render_as_string(hide_password=False)
|
||||
assert len(_alembic_heads(url)) == 1
|
||||
|
||||
def test_head_has_quota_schema(self, pg_engine):
|
||||
with pg_engine.connect() as conn:
|
||||
assert _alembic_version(conn) >= _0033
|
||||
assert _table_exists(conn, "quota_policies")
|
||||
assert _column_exists(conn, "token_usage", "cost")
|
||||
|
||||
def test_downgrade_drops_then_upgrade_restores(self, pg_engine):
|
||||
url = pg_engine.url.render_as_string(hide_password=False)
|
||||
_run_alembic(url, "downgrade", _0032)
|
||||
with pg_engine.connect() as conn:
|
||||
assert _alembic_version(conn) == _0032
|
||||
assert not _table_exists(conn, "quota_policies")
|
||||
assert not _column_exists(conn, "token_usage", "cost")
|
||||
_run_alembic(url, "upgrade", "head")
|
||||
with pg_engine.connect() as conn:
|
||||
assert _table_exists(conn, "quota_policies")
|
||||
assert _column_exists(conn, "token_usage", "cost")
|
||||
|
||||
def test_upgrade_tolerates_an_existing_cost_column(self, pg_engine):
|
||||
"""A database that already carries ``token_usage.cost`` upgrades cleanly."""
|
||||
url = pg_engine.url.render_as_string(hide_password=False)
|
||||
_run_alembic(url, "downgrade", _0032)
|
||||
with pg_engine.begin() as conn:
|
||||
conn.execute(text("ALTER TABLE token_usage ADD COLUMN cost NUMERIC(12,8) NOT NULL DEFAULT 0"))
|
||||
conn.execute(
|
||||
text(
|
||||
"INSERT INTO token_usage (user_id, prompt_tokens, generated_tokens, cost) "
|
||||
"VALUES ('u-mig33', 10, 1, 0.5)"
|
||||
)
|
||||
)
|
||||
_run_alembic(url, "upgrade", "head")
|
||||
with pg_engine.connect() as conn:
|
||||
cost = conn.execute(text("SELECT cost FROM token_usage WHERE user_id = 'u-mig33'")).scalar()
|
||||
assert float(cost) == 0.5
|
||||
|
||||
|
||||
class TestQuotaPolicyConstraints:
|
||||
def test_one_row_per_subject_and_bucket(self, pg_conn):
|
||||
_insert_policy(pg_conn, scope="instance", token_limit=10)
|
||||
with pytest.raises(IntegrityError):
|
||||
with pg_conn.begin_nested():
|
||||
_insert_policy(pg_conn, scope="instance", token_limit=20)
|
||||
_insert_policy(pg_conn, scope="instance", bucket="agent", token_limit=20)
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"values",
|
||||
[
|
||||
{"scope": "instance", "subject_id": "u1"},
|
||||
{"scope": "user"},
|
||||
{"scope": "user", "subject_id": "u1", "token_limit": 5, "token_unlimited": True},
|
||||
{"scope": "user", "subject_id": "u1", "cost_limit_usd": 5, "cost_unlimited": True},
|
||||
{"scope": "user", "subject_id": "u1", "token_limit": -1},
|
||||
{"scope": "user", "subject_id": "u1", "bucket": "nope"},
|
||||
{"scope": "org", "subject_id": "u1"},
|
||||
],
|
||||
)
|
||||
def test_invalid_rows_rejected(self, pg_conn, values):
|
||||
with pytest.raises(IntegrityError):
|
||||
with pg_conn.begin_nested():
|
||||
_insert_policy(pg_conn, **values)
|
||||
|
||||
def test_deleting_a_team_removes_its_policies(self, pg_conn):
|
||||
team_id = pg_conn.execute(
|
||||
text("INSERT INTO teams (name, slug, owner_id) VALUES ('Q', 'q-mig33', 'owner') RETURNING id")
|
||||
).scalar()
|
||||
_insert_policy(pg_conn, scope="team", subject_id=str(team_id), token_limit=10)
|
||||
_insert_policy(pg_conn, scope="user", subject_id=str(team_id), token_limit=10)
|
||||
pg_conn.execute(text("DELETE FROM teams WHERE id = :id"), {"id": team_id})
|
||||
scopes = pg_conn.execute(
|
||||
text("SELECT scope FROM quota_policies WHERE subject_id = :id"), {"id": str(team_id)}
|
||||
).scalars().all()
|
||||
assert scopes == ["user"]
|
||||
@@ -712,3 +712,56 @@ def test_decorator_omits_cache_bins_when_provider_reports_none(monkeypatch):
|
||||
assert row["cache_write_tokens"] is None
|
||||
assert llm.emitted[0]["cached_tokens"] is None
|
||||
assert llm.emitted[0]["cache_write_tokens"] is None
|
||||
|
||||
|
||||
def _persist_with_cost(monkeypatch, llm, cost_fn):
|
||||
from docsgpt.usage import _persist_call_usage
|
||||
|
||||
_install_fake_token_repo(monkeypatch)
|
||||
monkeypatch.setattr("docsgpt.usage.compute_cost_usd", cost_fn)
|
||||
_persist_call_usage(
|
||||
llm, {"prompt_tokens": 1000, "generated_tokens": 10, "cached_tokens": 400}
|
||||
)
|
||||
return _FakeTokenUsageRepo.last_instance.inserted[0]
|
||||
|
||||
|
||||
class _CostLLM:
|
||||
decoded_token = {"sub": "user_123"}
|
||||
user_api_key = None
|
||||
agent_id = None
|
||||
_canonical_model_id = "priced-model"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_persist_prices_the_call_by_canonical_model(monkeypatch):
|
||||
seen = {}
|
||||
|
||||
def cost_fn(model, prompt, generated, cached_tokens=None, cache_write_tokens=None):
|
||||
seen.update(model=model, prompt=prompt, generated=generated, cached=cached_tokens)
|
||||
return 0.0123
|
||||
|
||||
row = _persist_with_cost(monkeypatch, _CostLLM(), cost_fn)
|
||||
|
||||
assert row["cost"] == 0.0123
|
||||
assert seen == {"model": "priced-model", "prompt": 1000, "generated": 10, "cached": 400}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_persist_records_byom_calls_at_zero_cost(monkeypatch):
|
||||
llm = _CostLLM()
|
||||
llm._is_byom = True
|
||||
|
||||
row = _persist_with_cost(monkeypatch, llm, lambda *a, **k: 9.9)
|
||||
|
||||
assert row["cost"] == 0.0
|
||||
assert row["prompt_tokens"] == 1000
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_persist_keeps_the_row_when_pricing_fails(monkeypatch):
|
||||
def boom(*args, **kwargs):
|
||||
raise RuntimeError("registry unavailable")
|
||||
|
||||
row = _persist_with_cost(monkeypatch, _CostLLM(), boom)
|
||||
|
||||
assert row["cost"] == 0.0
|
||||
Reference in new issue
Block a user