mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-10-05 20:28:23 +00:00
update completion path
This commit is contained in:
@@ -287,16 +287,21 @@ def _sum_string_chars(node: Any) -> int:
|
|||||||
return 0
|
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:
|
def estimate_prompt_tokens(body: dict) -> int:
|
||||||
"""Conservatively estimate prompt tokens for the whole provider-bound body.
|
"""Conservatively estimate prompt tokens for the whole provider-bound body.
|
||||||
|
|
||||||
Unlike ``estimate_tokens`` (message text only), this walks every field, so
|
Every string counts, as do token IDs in legacy ``prompt`` arrays, so no
|
||||||
prompt weight hidden in tool schemas, tool-call arguments, ``system``, or
|
forwarded field can hide prompt weight and shrink its reservation.
|
||||||
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.
|
|
||||||
"""
|
"""
|
||||||
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]:
|
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")
|
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):
|
class TopupData(BaseModel):
|
||||||
"""Universal top-up data schema for Lightning Network invoices."""
|
"""Universal top-up data schema for Lightning Network invoices."""
|
||||||
|
|
||||||
@@ -653,17 +660,20 @@ class BaseUpstreamProvider:
|
|||||||
return "openrouter.ai" in (self.base_url or "")
|
return "openrouter.ai" in (self.base_url or "")
|
||||||
|
|
||||||
def prepare_request_body(
|
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:
|
) -> bytes | None:
|
||||||
"""Transform request body for provider-specific requirements.
|
"""Transform request body for provider-specific requirements.
|
||||||
|
|
||||||
Automatically transforms model names and, for streaming chat
|
Automatically transforms model names and opts streaming OpenAI
|
||||||
completions, opts the upstream into emitting per-chunk ``usage``
|
completion endpoints into emitting per-chunk ``usage`` so cost
|
||||||
so cost tracking can read real token counts instead of falling
|
tracking can read real token counts.
|
||||||
back to ``MaxCostData``.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
body: Original request body bytes
|
body: Original request body bytes
|
||||||
|
include_stream_usage: Opt a streaming completion into usage chunks
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Transformed request body bytes
|
Transformed request body bytes
|
||||||
@@ -706,14 +716,8 @@ class BaseUpstreamProvider:
|
|||||||
# OpenAI-compatible streaming responses omit ``usage`` unless the
|
# OpenAI-compatible streaming responses omit ``usage`` unless the
|
||||||
# request sets ``stream_options.include_usage = true``. Without it
|
# request sets ``stream_options.include_usage = true``. Without it
|
||||||
# we can't reconcile token counts at end of stream and must use
|
# we can't reconcile token counts at end of stream and must use
|
||||||
# the local request/response estimator. Discriminate
|
# the local request/response estimator.
|
||||||
# chat-completions-shaped requests by the ``messages`` field so we
|
if data.get("stream") is True and include_stream_usage:
|
||||||
# don't poke unrelated endpoints.
|
|
||||||
if (
|
|
||||||
data.get("stream") is True
|
|
||||||
and "messages" in data
|
|
||||||
and isinstance(data.get("messages"), list)
|
|
||||||
):
|
|
||||||
existing = data.get("stream_options")
|
existing = data.get("stream_options")
|
||||||
merged = dict(existing) if isinstance(existing, dict) else {}
|
merged = dict(existing) if isinstance(existing, dict) else {}
|
||||||
if merged.get("include_usage") is not True:
|
if merged.get("include_usage") is not True:
|
||||||
@@ -1022,6 +1026,7 @@ class BaseUpstreamProvider:
|
|||||||
reservation_snapshot: ReservationSnapshot | None = None,
|
reservation_snapshot: ReservationSnapshot | None = None,
|
||||||
client: httpx.AsyncClient | None = None,
|
client: httpx.AsyncClient | None = None,
|
||||||
request_body: bytes | None = None,
|
request_body: bytes | None = None,
|
||||||
|
legacy_completion: bool = False,
|
||||||
) -> StreamingResponse:
|
) -> StreamingResponse:
|
||||||
"""Handle streaming chat completion responses with token usage tracking and cost adjustment.
|
"""Handle streaming chat completion responses with token usage tracking and cost adjustment.
|
||||||
|
|
||||||
@@ -1060,6 +1065,7 @@ class BaseUpstreamProvider:
|
|||||||
last_model_seen: str | None = None
|
last_model_seen: str | None = None
|
||||||
usage_chunk_data: dict | None = None
|
usage_chunk_data: dict | None = None
|
||||||
done_seen: bool = False
|
done_seen: bool = False
|
||||||
|
stream_id: str | None = None
|
||||||
|
|
||||||
async def finalize_db_only() -> None:
|
async def finalize_db_only() -> None:
|
||||||
nonlocal usage_finalized
|
nonlocal usage_finalized
|
||||||
@@ -1118,7 +1124,7 @@ class BaseUpstreamProvider:
|
|||||||
* ``[DONE]`` is swallowed so it can be re-emitted exactly once at
|
* ``[DONE]`` is swallowed so it can be re-emitted exactly once at
|
||||||
end of stream.
|
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")
|
event = raw_event.strip(b"\r\n")
|
||||||
if not event:
|
if not event:
|
||||||
@@ -1171,9 +1177,12 @@ class BaseUpstreamProvider:
|
|||||||
or not isinstance(obj["id"], str)
|
or not isinstance(obj["id"], str)
|
||||||
or obj["id"] == "existing-id"
|
or obj["id"] == "existing-id"
|
||||||
):
|
):
|
||||||
if not hasattr(self, "_current_stream_id"):
|
if stream_id is None:
|
||||||
self._current_stream_id = f"chatcmpl-{uuid.uuid4()}"
|
id_prefix = "cmpl" if legacy_completion else "chatcmpl"
|
||||||
obj["id"] = self._current_stream_id
|
stream_id = f"{id_prefix}-{uuid.uuid4()}"
|
||||||
|
obj["id"] = stream_id
|
||||||
|
else:
|
||||||
|
stream_id = obj["id"]
|
||||||
if isinstance(obj.get("usage"), dict):
|
if isinstance(obj.get("usage"), dict):
|
||||||
# Capture usage for end-of-stream cost reconciliation.
|
# Capture usage for end-of-stream cost reconciliation.
|
||||||
# Some models (e.g. Gemini thinking models over the
|
# Some models (e.g. Gemini thinking models over the
|
||||||
@@ -1285,11 +1294,14 @@ class BaseUpstreamProvider:
|
|||||||
raise
|
raise
|
||||||
|
|
||||||
if usage_chunk_data is None:
|
if usage_chunk_data is None:
|
||||||
if not hasattr(self, "_current_stream_id"):
|
if stream_id is None:
|
||||||
self._current_stream_id = f"chatcmpl-{uuid.uuid4()}"
|
id_prefix = "cmpl" if legacy_completion else "chatcmpl"
|
||||||
|
stream_id = f"{id_prefix}-{uuid.uuid4()}"
|
||||||
usage_chunk_data = {
|
usage_chunk_data = {
|
||||||
"id": self._current_stream_id,
|
"id": stream_id,
|
||||||
"object": "chat.completion.chunk",
|
"object": "text_completion"
|
||||||
|
if legacy_completion
|
||||||
|
else "chat.completion.chunk",
|
||||||
"model": last_model_seen or "unknown",
|
"model": last_model_seen or "unknown",
|
||||||
"choices": [],
|
"choices": [],
|
||||||
"usage": {
|
"usage": {
|
||||||
@@ -1361,6 +1373,7 @@ class BaseUpstreamProvider:
|
|||||||
model_obj: Model | None = None,
|
model_obj: Model | None = None,
|
||||||
reservation_snapshot: ReservationSnapshot | None = None,
|
reservation_snapshot: ReservationSnapshot | None = None,
|
||||||
request_body: bytes | None = None,
|
request_body: bytes | None = None,
|
||||||
|
legacy_completion: bool = False,
|
||||||
) -> Response:
|
) -> Response:
|
||||||
"""Handle non-streaming chat completion responses with token usage tracking and cost adjustment.
|
"""Handle non-streaming chat completion responses with token usage tracking and cost adjustment.
|
||||||
|
|
||||||
@@ -1400,9 +1413,11 @@ class BaseUpstreamProvider:
|
|||||||
if requested_model:
|
if requested_model:
|
||||||
response_json["model"] = requested_model
|
response_json["model"] = requested_model
|
||||||
if "id" not in response_json or not isinstance(response_json["id"], str):
|
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 = MissingUsageEstimator(request_body, model_obj)
|
||||||
usage_estimator.observe(response_json)
|
usage_estimator.observe(response_json)
|
||||||
response_json["usage"] = usage_estimator.openai_response_data(
|
response_json["usage"] = usage_estimator.openai_response_data(
|
||||||
@@ -2923,6 +2938,7 @@ class BaseUpstreamProvider:
|
|||||||
Returns:
|
Returns:
|
||||||
Response or StreamingResponse from upstream with cost tracking
|
Response or StreamingResponse from upstream with cost tracking
|
||||||
"""
|
"""
|
||||||
|
completion_path = _openai_completion_path(path)
|
||||||
path = self.normalize_request_path(path, model_obj)
|
path = self.normalize_request_path(path, model_obj)
|
||||||
|
|
||||||
if (
|
if (
|
||||||
@@ -2951,7 +2967,11 @@ class BaseUpstreamProvider:
|
|||||||
(model_obj.forwarded_model_id or model_obj.id) if model_obj else None
|
(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(
|
logger.debug(
|
||||||
"Forwarding request to upstream",
|
"Forwarding request to upstream",
|
||||||
@@ -3048,7 +3068,7 @@ class BaseUpstreamProvider:
|
|||||||
return mapped_error
|
return mapped_error
|
||||||
|
|
||||||
if (
|
if (
|
||||||
path.endswith("chat/completions")
|
completion_path is not None
|
||||||
or path.endswith("embeddings")
|
or path.endswith("embeddings")
|
||||||
or path.endswith("messages")
|
or path.endswith("messages")
|
||||||
or path.endswith("messages/count_tokens")
|
or path.endswith("messages/count_tokens")
|
||||||
@@ -3117,7 +3137,7 @@ class BaseUpstreamProvider:
|
|||||||
await response.aclose()
|
await response.aclose()
|
||||||
await client.aclose()
|
await client.aclose()
|
||||||
|
|
||||||
if path.endswith("chat/completions"):
|
if completion_path is not None:
|
||||||
client_wants_streaming = False
|
client_wants_streaming = False
|
||||||
if request_body:
|
if request_body:
|
||||||
try:
|
try:
|
||||||
@@ -3163,6 +3183,7 @@ class BaseUpstreamProvider:
|
|||||||
reservation_snapshot=reservation_snapshot,
|
reservation_snapshot=reservation_snapshot,
|
||||||
client=client,
|
client=client,
|
||||||
request_body=request_body,
|
request_body=request_body,
|
||||||
|
legacy_completion=completion_path == "completions",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Handle both non-streaming chat completions and embeddings
|
# Handle both non-streaming chat completions and embeddings
|
||||||
@@ -3177,6 +3198,7 @@ class BaseUpstreamProvider:
|
|||||||
model_obj=model_obj,
|
model_obj=model_obj,
|
||||||
reservation_snapshot=reservation_snapshot,
|
reservation_snapshot=reservation_snapshot,
|
||||||
request_body=request_body,
|
request_body=request_body,
|
||||||
|
legacy_completion=completion_path == "completions",
|
||||||
)
|
)
|
||||||
finally:
|
finally:
|
||||||
await response.aclose()
|
await response.aclose()
|
||||||
@@ -4213,6 +4235,7 @@ class BaseUpstreamProvider:
|
|||||||
Returns:
|
Returns:
|
||||||
Response or StreamingResponse with refund if applicable
|
Response or StreamingResponse with refund if applicable
|
||||||
"""
|
"""
|
||||||
|
completion_path = _openai_completion_path(path)
|
||||||
if path.startswith("v1/"):
|
if path.startswith("v1/"):
|
||||||
path = path.replace("v1/", "")
|
path = path.replace("v1/", "")
|
||||||
|
|
||||||
@@ -4241,7 +4264,11 @@ class BaseUpstreamProvider:
|
|||||||
|
|
||||||
url = f"{self.base_url}/{path}"
|
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(
|
logger.debug(
|
||||||
"Forwarding request to upstream",
|
"Forwarding request to upstream",
|
||||||
@@ -4338,7 +4365,7 @@ class BaseUpstreamProvider:
|
|||||||
return error_response
|
return error_response
|
||||||
|
|
||||||
if (
|
if (
|
||||||
path.endswith("chat/completions")
|
completion_path is not None
|
||||||
or path.endswith("embeddings")
|
or path.endswith("embeddings")
|
||||||
or path.endswith("messages")
|
or path.endswith("messages")
|
||||||
or path.endswith("messages/count_tokens")
|
or path.endswith("messages/count_tokens")
|
||||||
|
|||||||
@@ -23,7 +23,11 @@ import litellm
|
|||||||
from fastapi.responses import Response
|
from fastapi.responses import Response
|
||||||
|
|
||||||
from ..core import get_logger
|
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
|
from ..payment.models import Model
|
||||||
|
|
||||||
logger = get_logger(__name__)
|
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 ""
|
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")
|
messages = body.get("messages")
|
||||||
if not isinstance(messages, list):
|
if not isinstance(messages, list):
|
||||||
messages = []
|
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")
|
system = body.get("system")
|
||||||
if isinstance(system, str) and system:
|
if isinstance(system, str) and system:
|
||||||
messages = [{"role": "system", "content": system}, *messages]
|
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
|
tools = body.get("tools") if isinstance(body.get("tools"), list) else None
|
||||||
|
|
||||||
return int(
|
return prompt_token_ids + int(
|
||||||
litellm.token_counter(
|
litellm.token_counter(
|
||||||
model=model,
|
model=model,
|
||||||
messages=messages,
|
messages=messages,
|
||||||
@@ -133,7 +154,9 @@ class MissingUsageEstimator:
|
|||||||
if self._input_tokens is not None:
|
if self._input_tokens is not None:
|
||||||
return self._input_tokens
|
return self._input_tokens
|
||||||
try:
|
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:
|
except Exception as exc:
|
||||||
self._input_tokens = estimate_prompt_tokens(self.body)
|
self._input_tokens = estimate_prompt_tokens(self.body)
|
||||||
logger.debug(
|
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:
|
def test_missing_usage_estimator_does_not_count_response_metadata() -> None:
|
||||||
estimator = MissingUsageEstimator(b"{}", None)
|
estimator = MissingUsageEstimator(b"{}", None)
|
||||||
estimator.observe(
|
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.
|
# 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({"tools": [{"data": hidden}]}) >= 1_000
|
||||||
assert estimate_prompt_tokens({"system": "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:
|
async def test_discount_cannot_be_dodged_by_hiding_prompt_in_tools() -> None:
|
||||||
|
|||||||
Reference in New Issue
Block a user