diff --git a/routstr/payment/helpers.py b/routstr/payment/helpers.py index 4d4696f2..4509649c 100644 --- a/routstr/payment/helpers.py +++ b/routstr/payment/helpers.py @@ -287,16 +287,21 @@ def _sum_string_chars(node: Any) -> int: return 0 +def _count_prompt_token_ids(node: Any) -> int: + if isinstance(node, int) and not isinstance(node, bool): + return 1 + if isinstance(node, list): + return sum(_count_prompt_token_ids(item) for item in node) + return 0 + + def estimate_prompt_tokens(body: dict) -> int: """Conservatively estimate prompt tokens for the whole provider-bound body. - Unlike ``estimate_tokens`` (message text only), this walks every field, so - prompt weight hidden in tool schemas, tool-call arguments, ``system``, or - any field forwarded in future cannot escape the reservation estimate. It - over-estimates rather than under-estimates: the result only shrinks a - discount against a reservation that settlement later refunds. + Every string counts, as do token IDs in legacy ``prompt`` arrays, so no + forwarded field can hide prompt weight and shrink its reservation. """ - return _sum_string_chars(body) // 3 + return _sum_string_chars(body) // 3 + _count_prompt_token_ids(body.get("prompt")) def _get_image_dimensions(image_data: bytes) -> tuple[int, int]: diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index f8b7fe3e..420832a1 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -248,6 +248,13 @@ def _is_json_content_type(content_type: str | None) -> bool: return main.startswith("application/") and main.endswith("+json") +def _openai_completion_path(path: str) -> str | None: + canonical = "/" + path.rstrip("/") + if canonical.endswith("/chat/completions"): + return "chat/completions" + return "completions" if canonical.endswith("/completions") else None + + class TopupData(BaseModel): """Universal top-up data schema for Lightning Network invoices.""" @@ -653,17 +660,20 @@ class BaseUpstreamProvider: return "openrouter.ai" in (self.base_url or "") def prepare_request_body( - self, body: bytes | None, model_obj: Model + self, + body: bytes | None, + model_obj: Model, + include_stream_usage: bool = False, ) -> bytes | None: """Transform request body for provider-specific requirements. - Automatically transforms model names and, for streaming chat - completions, opts the upstream into emitting per-chunk ``usage`` - so cost tracking can read real token counts instead of falling - back to ``MaxCostData``. + Automatically transforms model names and opts streaming OpenAI + completion endpoints into emitting per-chunk ``usage`` so cost + tracking can read real token counts. Args: body: Original request body bytes + include_stream_usage: Opt a streaming completion into usage chunks Returns: Transformed request body bytes @@ -706,14 +716,8 @@ class BaseUpstreamProvider: # OpenAI-compatible streaming responses omit ``usage`` unless the # request sets ``stream_options.include_usage = true``. Without it # we can't reconcile token counts at end of stream and must use - # the local request/response estimator. Discriminate - # chat-completions-shaped requests by the ``messages`` field so we - # don't poke unrelated endpoints. - if ( - data.get("stream") is True - and "messages" in data - and isinstance(data.get("messages"), list) - ): + # the local request/response estimator. + if data.get("stream") is True and include_stream_usage: existing = data.get("stream_options") merged = dict(existing) if isinstance(existing, dict) else {} if merged.get("include_usage") is not True: @@ -1022,6 +1026,7 @@ class BaseUpstreamProvider: reservation_snapshot: ReservationSnapshot | None = None, client: httpx.AsyncClient | None = None, request_body: bytes | None = None, + legacy_completion: bool = False, ) -> StreamingResponse: """Handle streaming chat completion responses with token usage tracking and cost adjustment. @@ -1060,6 +1065,7 @@ class BaseUpstreamProvider: last_model_seen: str | None = None usage_chunk_data: dict | None = None done_seen: bool = False + stream_id: str | None = None async def finalize_db_only() -> None: nonlocal usage_finalized @@ -1118,7 +1124,7 @@ class BaseUpstreamProvider: * ``[DONE]`` is swallowed so it can be re-emitted exactly once at end of stream. """ - nonlocal last_model_seen, usage_chunk_data, done_seen + nonlocal last_model_seen, usage_chunk_data, done_seen, stream_id event = raw_event.strip(b"\r\n") if not event: @@ -1171,9 +1177,12 @@ class BaseUpstreamProvider: or not isinstance(obj["id"], str) or obj["id"] == "existing-id" ): - if not hasattr(self, "_current_stream_id"): - self._current_stream_id = f"chatcmpl-{uuid.uuid4()}" - obj["id"] = self._current_stream_id + if stream_id is None: + id_prefix = "cmpl" if legacy_completion else "chatcmpl" + stream_id = f"{id_prefix}-{uuid.uuid4()}" + obj["id"] = stream_id + else: + stream_id = obj["id"] if isinstance(obj.get("usage"), dict): # Capture usage for end-of-stream cost reconciliation. # Some models (e.g. Gemini thinking models over the @@ -1285,11 +1294,14 @@ class BaseUpstreamProvider: raise if usage_chunk_data is None: - if not hasattr(self, "_current_stream_id"): - self._current_stream_id = f"chatcmpl-{uuid.uuid4()}" + if stream_id is None: + id_prefix = "cmpl" if legacy_completion else "chatcmpl" + stream_id = f"{id_prefix}-{uuid.uuid4()}" usage_chunk_data = { - "id": self._current_stream_id, - "object": "chat.completion.chunk", + "id": stream_id, + "object": "text_completion" + if legacy_completion + else "chat.completion.chunk", "model": last_model_seen or "unknown", "choices": [], "usage": { @@ -1361,6 +1373,7 @@ class BaseUpstreamProvider: model_obj: Model | None = None, reservation_snapshot: ReservationSnapshot | None = None, request_body: bytes | None = None, + legacy_completion: bool = False, ) -> Response: """Handle non-streaming chat completion responses with token usage tracking and cost adjustment. @@ -1400,9 +1413,11 @@ class BaseUpstreamProvider: if requested_model: response_json["model"] = requested_model if "id" not in response_json or not isinstance(response_json["id"], str): - response_json["id"] = f"chatcmpl-{uuid.uuid4()}" + prefix = "cmpl" if legacy_completion else "chatcmpl" + response_json["id"] = f"{prefix}-{uuid.uuid4()}" - if not isinstance(response_json.get("usage"), dict): + usage = response_json.get("usage") + if not isinstance(usage, dict) or not usage: usage_estimator = MissingUsageEstimator(request_body, model_obj) usage_estimator.observe(response_json) response_json["usage"] = usage_estimator.openai_response_data( @@ -2923,6 +2938,7 @@ class BaseUpstreamProvider: Returns: Response or StreamingResponse from upstream with cost tracking """ + completion_path = _openai_completion_path(path) path = self.normalize_request_path(path, model_obj) if ( @@ -2951,7 +2967,11 @@ class BaseUpstreamProvider: (model_obj.forwarded_model_id or model_obj.id) if model_obj else None ) - transformed_body = self.prepare_request_body(request_body, model_obj) + transformed_body = self.prepare_request_body( + request_body, + model_obj, + include_stream_usage=completion_path is not None, + ) logger.debug( "Forwarding request to upstream", @@ -3048,7 +3068,7 @@ class BaseUpstreamProvider: return mapped_error if ( - path.endswith("chat/completions") + completion_path is not None or path.endswith("embeddings") or path.endswith("messages") or path.endswith("messages/count_tokens") @@ -3117,7 +3137,7 @@ class BaseUpstreamProvider: await response.aclose() await client.aclose() - if path.endswith("chat/completions"): + if completion_path is not None: client_wants_streaming = False if request_body: try: @@ -3163,6 +3183,7 @@ class BaseUpstreamProvider: reservation_snapshot=reservation_snapshot, client=client, request_body=request_body, + legacy_completion=completion_path == "completions", ) # Handle both non-streaming chat completions and embeddings @@ -3177,6 +3198,7 @@ class BaseUpstreamProvider: model_obj=model_obj, reservation_snapshot=reservation_snapshot, request_body=request_body, + legacy_completion=completion_path == "completions", ) finally: await response.aclose() @@ -4213,6 +4235,7 @@ class BaseUpstreamProvider: Returns: Response or StreamingResponse with refund if applicable """ + completion_path = _openai_completion_path(path) if path.startswith("v1/"): path = path.replace("v1/", "") @@ -4241,7 +4264,11 @@ class BaseUpstreamProvider: url = f"{self.base_url}/{path}" - transformed_body = self.prepare_request_body(request_body, model_obj) + transformed_body = self.prepare_request_body( + request_body, + model_obj, + include_stream_usage=completion_path is not None, + ) logger.debug( "Forwarding request to upstream", @@ -4338,7 +4365,7 @@ class BaseUpstreamProvider: return error_response if ( - path.endswith("chat/completions") + completion_path is not None or path.endswith("embeddings") or path.endswith("messages") or path.endswith("messages/count_tokens") diff --git a/routstr/upstream/count_tokens.py b/routstr/upstream/count_tokens.py index c79d11d9..8ebeb6d9 100644 --- a/routstr/upstream/count_tokens.py +++ b/routstr/upstream/count_tokens.py @@ -23,7 +23,11 @@ import litellm from fastapi.responses import Response from ..core import get_logger -from ..payment.helpers import estimate_prompt_tokens, estimate_tokens +from ..payment.helpers import ( + _count_prompt_token_ids, + estimate_prompt_tokens, + estimate_tokens, +) from ..payment.models import Model logger = get_logger(__name__) @@ -46,11 +50,28 @@ def _model_name(model_obj: Model | None, body: dict[str, Any]) -> str: return body_model if isinstance(body_model, str) else "" -def _count_with_litellm(model: str, body: dict[str, Any]) -> int: +def _count_with_litellm( + model: str, body: dict[str, Any], include_legacy_prompt: bool = False +) -> int: messages = body.get("messages") if not isinstance(messages, list): messages = [] + prompt_token_ids = 0 + if include_legacy_prompt: + prompt = body.get("prompt") + if isinstance(prompt, str): + prompt_texts = [prompt] + elif isinstance(prompt, list): + prompt_texts = [item for item in prompt if isinstance(item, str)] + else: + prompt_texts = [] + prompt_token_ids = _count_prompt_token_ids(prompt) + messages = [ + *({"role": "user", "content": text} for text in prompt_texts if text), + *messages, + ] + system = body.get("system") if isinstance(system, str) and system: messages = [{"role": "system", "content": system}, *messages] @@ -65,7 +86,7 @@ def _count_with_litellm(model: str, body: dict[str, Any]) -> int: tools = body.get("tools") if isinstance(body.get("tools"), list) else None - return int( + return prompt_token_ids + int( litellm.token_counter( model=model, messages=messages, @@ -133,7 +154,9 @@ class MissingUsageEstimator: if self._input_tokens is not None: return self._input_tokens try: - self._input_tokens = _count_with_litellm(self.model_name, self.body) + self._input_tokens = _count_with_litellm( + self.model_name, self.body, include_legacy_prompt=True + ) except Exception as exc: self._input_tokens = estimate_prompt_tokens(self.body) logger.debug( diff --git a/tests/unit/test_completions_billing.py b/tests/unit/test_completions_billing.py new file mode 100644 index 00000000..6a512136 --- /dev/null +++ b/tests/unit/test_completions_billing.py @@ -0,0 +1,384 @@ +import json +from typing import Any +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest +from fastapi.responses import Response +from sqlalchemy.ext.asyncio import AsyncEngine, create_async_engine +from sqlmodel import SQLModel +from sqlmodel.ext.asyncio.session import AsyncSession + +import routstr.auth as auth_module +from routstr.auth import ReservationSnapshot, get_reservation_snapshot, pay_for_request +from routstr.core.db import ApiKey, ReservationRelease +from routstr.payment.models import Architecture, Model, Pricing +from routstr.upstream.base import BaseUpstreamProvider + +BALANCE = 100_000 +RESERVED = 5_000 +# 1 msat per prompt token, 2 msat per completion token. +MODEL = Model( + id="glm-test", + name="glm-test", + created=0, + description="", + context_length=64_000, + architecture=Architecture( + modality="text->text", + input_modalities=["text"], + output_modalities=["text"], + tokenizer="Other", + instruct_type=None, + ), + pricing=Pricing(prompt=0.001, completion=0.002), + sats_pricing=Pricing(prompt=0.001, completion=0.002), +) +USAGE = {"prompt_tokens": 400, "completion_tokens": 100, "total_tokens": 500} +EXPECTED_MSATS = 400 * 1 + 100 * 2 + +COMPLETION_BODY = { + "model": MODEL.id, + "prompt": "Once upon a time, in a land far away, " * 20, + "max_tokens": 100, +} +CHAT_BODY = {"model": MODEL.id, "messages": [{"role": "user", "content": "hi"}]} + +COMPLETION_JSON = { + "id": "cmpl-1", + "object": "text_completion", + "model": MODEL.id, + "choices": [{"text": " there was", "index": 0, "finish_reason": "stop"}], + "usage": USAGE, +} +COMPLETION_CHUNKS = [ + { + "id": "cmpl-1", + "object": "text_completion", + "model": MODEL.id, + "choices": [{"text": " there", "index": 0, "finish_reason": None}], + }, + { + "id": "cmpl-1", + "object": "text_completion", + "model": MODEL.id, + "choices": [{"text": " was", "index": 0, "finish_reason": "stop"}], + }, +] +USAGE_CHUNK = { + "id": "cmpl-1", + "object": "text_completion", + "model": MODEL.id, + "choices": [], + "usage": USAGE, +} + + +def _sse(chunks: list[dict]) -> bytes: + body = b"".join(b"data: " + json.dumps(c).encode() + b"\n\n" for c in chunks) + return body + b"data: [DONE]\n\n" + + +@pytest.fixture(autouse=True) +def patch_sats_usd_price() -> Any: + with patch("routstr.payment.cost_calculation.sats_usd_price", return_value=5.0e-4): + yield + + +async def _engine() -> AsyncEngine: + engine = create_async_engine("sqlite+aiosqlite://") + async with engine.begin() as connection: + await connection.run_sync(SQLModel.metadata.create_all) + return engine + + +def _upstream(content: bytes, content_type: str) -> httpx.Response: + return httpx.Response( + 200, + content=content, + headers={"content-type": content_type}, + request=httpx.Request("POST", "http://upstream"), + ) + + +async def _drain(response: Any) -> bytes: + body = b"" + if hasattr(response, "body_iterator"): + async for chunk in response.body_iterator: + body += chunk if isinstance(chunk, bytes) else chunk.encode() + else: + body = response.body + return body + + +async def _forward( + engine: AsyncEngine, + path: str, + body: dict, + upstream: httpx.Response, +) -> tuple[bytes, ReservationSnapshot, AsyncMock]: + """Reserve, forward through the real ``forward_request`` and settle.""" + provider = BaseUpstreamProvider( + base_url="http://upstream", api_key="k", provider_fee=1.0 + ) + request = MagicMock() + request.method = "POST" + request.query_params = {} + send = AsyncMock(return_value=upstream) + + async with AsyncSession(engine, expire_on_commit=False) as session: + key = ApiKey(hashed_key="key", balance=BALANCE) + session.add(key) + await session.commit() + await pay_for_request(key, RESERVED, session) + snapshot = await get_reservation_snapshot(key, session) + + with ( + patch("httpx.AsyncClient.send", send), + patch( + "routstr.upstream.base.create_session", + side_effect=lambda: AsyncSession(engine, expire_on_commit=False), + ), + patch( + "routstr.upstream.base.adjust_payment_for_tokens", + auth_module.adjust_payment_for_tokens, + ), + ): + response = await provider.forward_request( + request, + path, + {}, + json.dumps(body).encode(), + key, + RESERVED, + session, + MODEL, + snapshot, + ) + out = await _drain(response) + return out, snapshot, send + + +async def _ledger( + engine: AsyncEngine, snapshot: ReservationSnapshot +) -> tuple[int, int, int, str | None]: + async with AsyncSession(engine, expire_on_commit=False) as session: + key = await session.get(ApiKey, snapshot.key_hash) + record = await session.get(ReservationRelease, snapshot.release_id) + assert key is not None + return ( + key.balance, + key.total_spent, + key.reserved_balance, + record.status if record else None, + ) + + +def _sse_objects(out: bytes) -> list[dict]: + objs = [] + for line in out.split(b"\n"): + if line.startswith(b"data: ") and line[6:].strip() != b"[DONE]": + objs.append(json.loads(line[6:])) + return objs + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "path", + [ + "completions", + "v1/completions", + "v1/completions/", + "openai/v1/completions", + ], +) +async def test_non_streaming_completion_with_usage_is_charged(path: str) -> None: + engine = await _engine() + out, snapshot, _ = await _forward( + engine, + path, + COMPLETION_BODY, + _upstream(json.dumps(COMPLETION_JSON).encode(), "application/json"), + ) + + balance, spent, reserved, status = await _ledger(engine, snapshot) + assert (balance, spent, reserved, status) == ( + BALANCE - EXPECTED_MSATS, + EXPECTED_MSATS, + 0, + "charged", + ) + body = json.loads(out) + assert body["object"] == "text_completion" + assert body["usage"]["cost"]["total_msats"] == EXPECTED_MSATS + await engine.dispose() + + +@pytest.mark.asyncio +async def test_streaming_completion_with_final_usage_is_charged() -> None: + engine = await _engine() + out, snapshot, send = await _forward( + engine, + "v1/completions", + {**COMPLETION_BODY, "stream": True}, + _upstream(_sse([*COMPLETION_CHUNKS, USAGE_CHUNK]), "text/event-stream"), + ) + + balance, spent, reserved, status = await _ledger(engine, snapshot) + assert (balance, spent, reserved, status) == ( + BALANCE - EXPECTED_MSATS, + EXPECTED_MSATS, + 0, + "charged", + ) + + forwarded = json.loads(send.call_args.args[0].content) + assert forwarded["stream_options"] == {"include_usage": True} + + objs = _sse_objects(out) + assert [o["choices"][0]["text"] for o in objs if o["choices"]] == [ + " there", + " was", + ] + assert objs[-1]["object"] == "text_completion" + assert objs[-1]["usage"]["cost"]["total_msats"] == EXPECTED_MSATS + assert out.endswith(b"data: [DONE]\n\n") + await engine.dispose() + + +@pytest.mark.asyncio +async def test_non_streaming_completion_with_empty_usage_is_estimated() -> None: + """Missing usage is estimated from ``prompt`` and ``text``, never free.""" + engine = await _engine() + no_usage = {k: v for k, v in COMPLETION_JSON.items() if k != "id"} + no_usage["usage"] = {} + out, snapshot, _ = await _forward( + engine, + "v1/completions", + COMPLETION_BODY, + _upstream(json.dumps(no_usage).encode(), "application/json"), + ) + + balance, spent, reserved, status = await _ledger(engine, snapshot) + assert 0 < spent <= RESERVED + assert (BALANCE - balance, reserved, status) == (spent, 0, "charged") + body = json.loads(out) + usage = body["usage"] + assert body["id"].startswith("cmpl-") + assert usage["estimated"] is True + assert usage["prompt_tokens"] > 100 + assert usage["completion_tokens"] > 0 + await engine.dispose() + + +@pytest.mark.asyncio +async def test_streaming_completion_without_usage_is_estimated() -> None: + engine = await _engine() + out, snapshot, _ = await _forward( + engine, + "v1/completions", + {**COMPLETION_BODY, "stream": True}, + _upstream(_sse(COMPLETION_CHUNKS), "text/event-stream"), + ) + + balance, spent, reserved, status = await _ledger(engine, snapshot) + assert 0 < spent <= RESERVED + assert (BALANCE - balance, reserved, status) == (spent, 0, "charged") + trailer = _sse_objects(out)[-1] + assert trailer["id"] == "cmpl-1" + assert trailer["object"] == "text_completion" + assert trailer["usage"]["prompt_tokens"] > 100 + assert trailer["usage"]["cost"]["total_msats"] == spent + await engine.dispose() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("path", ["v1/chat/completions/", "openai/v1/chat/completions"]) +async def test_chat_completion_aliases_are_charged(path: str) -> None: + engine = await _engine() + chat_json = { + "id": "chatcmpl-1", + "object": "chat.completion", + "model": MODEL.id, + "choices": [ + { + "message": {"role": "assistant", "content": "hi"}, + "index": 0, + "finish_reason": "stop", + } + ], + "usage": USAGE, + } + out, snapshot, _ = await _forward( + engine, + path, + CHAT_BODY, + _upstream(json.dumps(chat_json).encode(), "application/json"), + ) + + balance, spent, reserved, status = await _ledger(engine, snapshot) + assert (balance, spent, reserved, status) == ( + BALANCE - EXPECTED_MSATS, + EXPECTED_MSATS, + 0, + "charged", + ) + assert json.loads(out)["usage"]["cost"]["total_msats"] == EXPECTED_MSATS + await engine.dispose() + + +def test_stream_usage_option_is_scoped_to_completion_endpoints() -> None: + provider = BaseUpstreamProvider(base_url="http://upstream", api_key="k") + body = json.dumps({"prompt": "draw this", "stream": True}).encode() + + assert provider.prepare_request_body(body, MODEL) == body + + +async def _forward_x_cashu( + path: str, body: dict, upstream: httpx.Response +) -> tuple[AsyncMock, AsyncMock]: + """Run ``forward_x_cashu_request`` with the settlement handler stubbed out.""" + provider = BaseUpstreamProvider( + base_url="http://upstream", api_key="k", provider_fee=1.0 + ) + request = MagicMock() + request.method = "POST" + request.query_params = {} + request.state.request_id = "req-1" + request.body = AsyncMock(return_value=json.dumps(body).encode()) + send = AsyncMock(return_value=upstream) + settle = AsyncMock(return_value=Response(content=b"{}", status_code=200)) + + with ( + patch("httpx.AsyncClient.send", send), + patch.object(provider, "handle_x_cashu_chat_completion", settle), + ): + await provider.forward_x_cashu_request( + request, path, {}, 10, "sat", RESERVED, MODEL + ) + return settle, send + + +@pytest.mark.asyncio +@pytest.mark.parametrize("path", ["completions", "v1/completions", "v1/completions/"]) +async def test_x_cashu_legacy_completion_is_settled(path: str) -> None: + """Legacy completions must reach refund settlement, not raw passthrough.""" + settle, _ = await _forward_x_cashu( + path, + COMPLETION_BODY, + _upstream(json.dumps(COMPLETION_JSON).encode(), "application/json"), + ) + + assert settle.await_count == 1 + + +@pytest.mark.asyncio +async def test_x_cashu_streaming_completion_requests_usage() -> None: + _, send = await _forward_x_cashu( + "v1/completions", + {**COMPLETION_BODY, "stream": True}, + _upstream(_sse([*COMPLETION_CHUNKS, USAGE_CHUNK]), "text/event-stream"), + ) + + forwarded = json.loads(send.call_args.args[0].content) + assert forwarded["stream_options"] == {"include_usage": True} diff --git a/tests/unit/test_count_tokens_local.py b/tests/unit/test_count_tokens_local.py index ebd23cb2..85a8ec72 100644 --- a/tests/unit/test_count_tokens_local.py +++ b/tests/unit/test_count_tokens_local.py @@ -229,6 +229,15 @@ def test_missing_usage_estimator_openai_dialect() -> None: } +def test_missing_usage_estimator_counts_legacy_token_prompt() -> None: + model = _make_model() + request_body = _body({"model": model.id, "prompt": [[1, 2], [3, 4, 5]]}) + + usage = MissingUsageEstimator(request_body, model).openai_response_data()["usage"] + + assert usage["prompt_tokens"] >= 5 + + def test_missing_usage_estimator_does_not_count_response_metadata() -> None: estimator = MissingUsageEstimator(b"{}", None) estimator.observe( diff --git a/tests/unit/test_payment_helpers.py b/tests/unit/test_payment_helpers.py index e5dbf136..2eab19c1 100644 --- a/tests/unit/test_payment_helpers.py +++ b/tests/unit/test_payment_helpers.py @@ -183,6 +183,32 @@ def test_estimate_prompt_tokens_counts_every_string_in_the_body() -> None: # value prefix can buy a discount, so both still count in full. assert estimate_prompt_tokens({"tools": [{"data": hidden}]}) >= 1_000 assert estimate_prompt_tokens({"system": "data:" + hidden}) >= 1_000 + assert estimate_prompt_tokens({"prompt": [[1, 2], [3, 4, 5]]}) >= 5 + + +async def test_discount_counts_legacy_token_id_prompt() -> None: + from routstr.payment.helpers import calculate_discounted_max_cost + + pricing = Mock() + pricing.prompt = 0.001 + pricing.completion = 0.0 + pricing.max_prompt_cost = 50.0 + pricing.max_completion_cost = 0.0 + + model_obj = Mock() + model_obj.sats_pricing = pricing + model_obj.top_provider = None + model_obj.context_length = None + + body = {"model": "test-model", "prompt": list(range(50_000)), "max_tokens": 0} + with ( + patch.object(settings, "fixed_pricing", False), + patch.object(settings, "tolerance_percentage", 0), + patch.object(settings, "min_request_msat", 1000), + ): + cost = await calculate_discounted_max_cost(50_000, body, model_obj) + + assert cost == 50_000 async def test_discount_cannot_be_dodged_by_hiding_prompt_in_tools() -> None: