Compare commits

...
Author SHA1 Message Date
9qeklajc 30db582321 add model id to validation log 2026-04-08 00:36:05 +02:00
9qeklajcandGitHub 5ef499e2ae Merge pull request #443 from Routstr/437-clean-up-logging
do not log client host address
2026-04-07 00:44:07 +02:00
9qeklajc c7c802c610 do not log client host address 2026-04-07 00:23:53 +02:00
9qeklajcandGitHub 55e240d92a Merge pull request #435 from Routstr/add-missing-sat-cost-in-x-cashu
add sat cost to x-cashu response
2026-04-05 01:35:17 +02:00
9qeklajc 7da4ad3818 add sat cost to x-cashu response 2026-04-05 01:20:58 +02:00
9qeklajcandGitHub 9e1934bfda Merge pull request #434 from Routstr/opencode-display-usage-correctly
Opencode display usage correctly
2026-04-05 00:34:22 +02:00
9qeklajcandGitHub d686e0e851 Merge pull request #433 from Routstr/update-x-cashu-swept-default-to-one-week
x-cashu swept default to one week
2026-04-05 00:23:37 +02:00
9qeklajc 7708ed1c8b x-cashu swept default to one week 2026-04-05 00:21:53 +02:00
5 changed files with 257 additions and 14 deletions
-8
View File
@@ -38,11 +38,6 @@ class LoggingMiddleware(BaseHTTPMiddleware):
except Exception:
pass
# Extract request info
client_host = None
if request.client:
client_host = request.client.host
# Log incoming request
logger.info(
"Incoming request",
@@ -51,7 +46,6 @@ class LoggingMiddleware(BaseHTTPMiddleware):
"method": request.method,
"path": request.url.path,
"query_params": dict(request.query_params),
"client_host": client_host,
"headers": {
k: v
for k, v in request.headers.items()
@@ -100,7 +94,6 @@ class LoggingMiddleware(BaseHTTPMiddleware):
"path": request.url.path,
"status_code": response.status_code,
"duration_ms": round(duration * 1000, 2),
"client_host": client_host,
},
)
if hasattr(response, "headers"):
@@ -120,7 +113,6 @@ class LoggingMiddleware(BaseHTTPMiddleware):
"method": request.method,
"path": request.url.path,
"duration_ms": round(duration * 1000, 2),
"client_host": client_host,
"error": str(e),
"error_type": type(e).__name__,
},
+1 -1
View File
@@ -74,7 +74,7 @@ class Settings(BaseSettings):
enable_pricing_refresh: bool = Field(default=True, env="ENABLE_PRICING_REFRESH")
enable_models_refresh: bool = Field(default=True, env="ENABLE_MODELS_REFRESH")
refund_cache_ttl_seconds: int = Field(default=3600, env="REFUND_CACHE_TTL_SECONDS")
refund_sweep_ttl_seconds: int = Field(default=86400, env="REFUND_SWEEP_TTL_SECONDS")
refund_sweep_ttl_seconds: int = Field(default=604800, env="REFUND_SWEEP_TTL_SECONDS")
# Logging
log_level: str = Field(default="INFO", env="LOG_LEVEL")
+10 -3
View File
@@ -207,7 +207,7 @@ async def proxy(
elif auth := headers.get("authorization", None):
key = await get_bearer_token_key(
headers, path, session, auth, max_cost_for_model
headers, path, session, auth, max_cost_for_model, model_id
)
else:
@@ -387,7 +387,12 @@ async def proxy(
async def get_bearer_token_key(
headers: dict, path: str, session: AsyncSession, auth: str, min_cost: int = 0
headers: dict,
path: str,
session: AsyncSession,
auth: str,
min_cost: int = 0,
model_id: str = "unknown",
) -> ApiKey:
"""Handle bearer token authentication proxy requests."""
parts = auth.split()
@@ -457,11 +462,13 @@ async def get_bearer_token_key(
except Exception as e:
key_preview = bearer_key[:20] + "..." if len(bearer_key) > 20 else bearer_key
logger.error(
f"Bearer token validation failed: {type(e).__name__}: {e} path={path} key={key_preview!r}",
f"Bearer token validation failed: {type(e).__name__}: {e} path={path} model={model_id!r} min_cost={min_cost} key={key_preview!r}",
extra={
"error": str(e),
"error_type": type(e).__name__,
"path": path,
"model_id": model_id,
"min_cost_msat": min_cost,
"bearer_key_preview": key_preview,
},
)
+30 -2
View File
@@ -2152,6 +2152,17 @@ class BaseUpstreamProvider:
},
)
if cost_data:
for i, line in enumerate(lines):
if line.startswith("data: "):
try:
data_json = json.loads(line[6:])
if "usage" in data_json and data_json["usage"]:
data_json["usage"]["cost_sats"] = cost_data.total_msats // 1000
lines[i] = "data: " + json.dumps(data_json)
except json.JSONDecodeError:
pass
async def generate() -> AsyncGenerator[bytes, None]:
for line in lines:
yield (line + "\n").encode("utf-8")
@@ -2196,6 +2207,9 @@ class BaseUpstreamProvider:
response_json = json.loads(content_str)
cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model)
if cost_data and "usage" in response_json:
response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000
if not cost_data:
logger.error(
"Failed to calculate cost for response",
@@ -2262,7 +2276,7 @@ class BaseUpstreamProvider:
)
return Response(
content=content_str,
content=json.dumps(response_json),
status_code=response.status_code,
headers=response_headers,
media_type="application/json",
@@ -3067,6 +3081,17 @@ class BaseUpstreamProvider:
},
)
if cost_data:
for i, line in enumerate(lines):
if line.startswith("data: "):
try:
data_json = json.loads(line[6:])
if "usage" in data_json and data_json["usage"]:
data_json["usage"]["cost_sats"] = cost_data.total_msats // 1000
lines[i] = "data: " + json.dumps(data_json)
except json.JSONDecodeError:
pass
async def generate() -> AsyncGenerator[bytes, None]:
for line in lines:
yield (line + "\\n").encode("utf-8")
@@ -3099,6 +3124,9 @@ class BaseUpstreamProvider:
response_json = json.loads(content_str)
cost_data = await self.get_x_cashu_cost(response_json, max_cost_for_model)
if cost_data and "usage" in response_json:
response_json["usage"]["cost_sats"] = cost_data.total_msats // 1000
if not cost_data:
logger.error(
"Failed to calculate cost for Responses API response",
@@ -3165,7 +3193,7 @@ class BaseUpstreamProvider:
)
return Response(
content=content_str,
content=json.dumps(response_json),
status_code=response.status_code,
headers=response_headers,
media_type="application/json",
+216
View File
@@ -0,0 +1,216 @@
import json
import os
from unittest.mock import AsyncMock, patch
import httpx
import pytest
os.environ.setdefault("UPSTREAM_BASE_URL", "http://test")
os.environ.setdefault("UPSTREAM_API_KEY", "test")
from routstr.payment.cost_calculation import CostData # noqa: E402
from routstr.upstream.base import BaseUpstreamProvider # noqa: E402
def _make_provider() -> BaseUpstreamProvider:
return BaseUpstreamProvider(base_url="http://test", api_key="test-key")
def _make_httpx_response(status_code: int = 200) -> httpx.Response:
return httpx.Response(status_code, headers={})
def _make_cost_data(total_msats: int = 5000) -> CostData:
return CostData(
base_msats=0,
input_msats=3000,
output_msats=2000,
total_msats=total_msats,
total_usd=0.00025,
input_tokens=100,
output_tokens=50,
)
# ---------------------------------------------------------------------------
# Non-streaming (chat completions)
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_non_streaming_includes_cost_sats() -> None:
provider = _make_provider()
cost_data = _make_cost_data(total_msats=5000)
response_body = {
"model": "gpt-4o",
"usage": {
"prompt_tokens": 100,
"completion_tokens": 50,
"total_tokens": 150,
"cost": 0.00025,
},
}
content_str = json.dumps(response_body)
httpx_response = _make_httpx_response()
with (
patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=cost_data)),
patch.object(provider, "send_refund", new=AsyncMock(return_value="cashuA_refund_token")),
):
response = await provider.handle_x_cashu_non_streaming_response(
content_str=content_str,
response=httpx_response,
amount=10000,
unit="msat",
max_cost_for_model=10000,
mint=None,
payment_token_hash=None,
)
body = json.loads(response.body)
assert "cost_sats" in body["usage"]
assert body["usage"]["cost_sats"] == 5 # 5000 msats // 1000
@pytest.mark.asyncio
async def test_non_streaming_cost_sats_value_rounds_down() -> None:
provider = _make_provider()
cost_data = _make_cost_data(total_msats=1999)
response_body = {"model": "gpt-4o", "usage": {"prompt_tokens": 10}}
content_str = json.dumps(response_body)
with (
patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=cost_data)),
patch.object(provider, "send_refund", new=AsyncMock(return_value="cashuA_refund_token")),
):
response = await provider.handle_x_cashu_non_streaming_response(
content_str=content_str,
response=_make_httpx_response(),
amount=10000,
unit="msat",
max_cost_for_model=10000,
)
body = json.loads(response.body)
assert body["usage"]["cost_sats"] == 1 # 1999 // 1000
@pytest.mark.asyncio
async def test_non_streaming_preserves_existing_usage_fields() -> None:
provider = _make_provider()
cost_data = _make_cost_data(total_msats=3000)
response_body = {
"model": "gpt-4o",
"usage": {
"prompt_tokens": 100,
"completion_tokens": 50,
"total_tokens": 150,
"cost": 0.00015,
},
}
with (
patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=cost_data)),
patch.object(provider, "send_refund", new=AsyncMock(return_value="cashuA_refund_token")),
):
response = await provider.handle_x_cashu_non_streaming_response(
content_str=json.dumps(response_body),
response=_make_httpx_response(),
amount=10000,
unit="msat",
max_cost_for_model=10000,
)
body = json.loads(response.body)
usage = body["usage"]
assert usage["prompt_tokens"] == 100
assert usage["completion_tokens"] == 50
assert usage["total_tokens"] == 150
assert usage["cost"] == 0.00015
assert usage["cost_sats"] == 3
# ---------------------------------------------------------------------------
# Streaming (chat completions)
# ---------------------------------------------------------------------------
async def _collect_streaming(response: object) -> list[str]:
chunks: list[str] = []
async for chunk in response.body_iterator: # type: ignore[attr-defined]
if isinstance(chunk, bytes):
chunks.append(chunk.decode("utf-8"))
else:
chunks.append(str(chunk))
return chunks
@pytest.mark.asyncio
async def test_streaming_includes_cost_sats_in_usage_chunk() -> None:
provider = _make_provider()
cost_data = _make_cost_data(total_msats=7000)
usage_chunk = {
"id": "chatcmpl-123",
"model": "gpt-4o",
"usage": {"prompt_tokens": 100, "completion_tokens": 50, "total_tokens": 150},
}
content_str = "\n".join([
'data: {"id":"chatcmpl-123","model":"gpt-4o","choices":[]}',
f"data: {json.dumps(usage_chunk)}",
"data: [DONE]",
])
with patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=cost_data)):
response = await provider.handle_x_cashu_streaming_response(
content_str=content_str,
response=_make_httpx_response(),
amount=10000,
unit="msat",
max_cost_for_model=10000,
mint=None,
payment_token_hash=None,
)
chunks = await _collect_streaming(response)
full_output = "".join(chunks)
usage_line = next(
line for line in full_output.split("\n") if '"usage"' in line and "cost_sats" in line
)
data_json = json.loads(usage_line.lstrip("data: ").strip())
assert data_json["usage"]["cost_sats"] == 7 # 7000 // 1000
@pytest.mark.asyncio
async def test_streaming_non_usage_chunks_unmodified() -> None:
provider = _make_provider()
cost_data = _make_cost_data(total_msats=2000)
regular_chunk = {"id": "chatcmpl-123", "model": "gpt-4o", "choices": [{"delta": {"content": "hi"}}]}
usage_chunk = {"id": "chatcmpl-123", "model": "gpt-4o", "usage": {"prompt_tokens": 10}}
content_str = "\n".join([
f"data: {json.dumps(regular_chunk)}",
f"data: {json.dumps(usage_chunk)}",
"data: [DONE]",
])
with patch.object(provider, "get_x_cashu_cost", new=AsyncMock(return_value=cost_data)):
response = await provider.handle_x_cashu_streaming_response(
content_str=content_str,
response=_make_httpx_response(),
amount=10000,
unit="msat",
max_cost_for_model=10000,
)
chunks = await _collect_streaming(response)
lines = [
line for line in "".join(chunks).split("\n")
if line.startswith("data: ") and line != "data: [DONE]"
]
regular_line_data = json.loads(lines[0][6:])
# regular chunk should not have cost_sats injected
assert "cost_sats" not in regular_line_data.get("usage", {})