mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 12:28:22 +00:00
Merge pull request #706 from Routstr/optimize-completion
update completion path
This commit is contained in:
@@ -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]:
|
||||
|
||||
+55
-28
@@ -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")
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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}
|
||||
@@ -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(
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user