From 8c0ac499ef4234852572ac14a2dd7fc50f19eb4f Mon Sep 17 00:00:00 2001 From: 9qeklajc Date: Mon, 27 Apr 2026 23:48:32 +0200 Subject: [PATCH] make sure to always emit sats cost --- routstr/upstream/base.py | 284 +++++++++++++++---------- tests/unit/test_stream_id_injection.py | 10 +- 2 files changed, 186 insertions(+), 108 deletions(-) diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 458e5dfd..6176c509 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -604,57 +604,83 @@ class BaseUpstreamProvider: ) yield prefix + part - # Stream finished, process usage if found - if usage_chunk_data: - async with create_session() as session: - fresh_key = await session.get(key.__class__, key.hashed_key) - if fresh_key: - try: - cost_data = await adjust_payment_for_tokens( - fresh_key, - usage_chunk_data, - session, - max_cost_for_model, - ) - remaining_balance_msats = fresh_key.balance - # Merge cost into usage - usage_chunk_data["usage"]["cost"] = cost_data.get( - "total_usd", 0.0 - ) - usage_chunk_data["usage"]["cost_sats"] = ( - cost_data.get("total_msats", 0) // 1000 - ) - usage_chunk_data["usage"]["remaining_balance_msats"] = ( - remaining_balance_msats - ) - # Keep detailed cost in metadata - usage_chunk_data["metadata"] = usage_chunk_data.get( - "metadata", {} - ) - usage_chunk_data["metadata"]["routstr"] = { - "cost": cost_data + async with create_session() as session: + fresh_key = await session.get(key.__class__, key.hashed_key) + if fresh_key: + cost_data: dict + try: + adjustment_input = ( + usage_chunk_data + if usage_chunk_data is not None + else { + "model": last_model_seen or "unknown", + "usage": None, } - usage_chunk_data["metadata"]["routstr"]["cost"][ - "sats_cost" - ] = cost_data.get("total_msats", 0) // 1000 - usage_chunk_data["metadata"]["routstr"]["cost"][ - "remaining_balance_msats" - ] = remaining_balance_msats - yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode() - usage_finalized = True - except Exception as e: - logger.exception( - "Error during usage finalization", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "error": str(e), - }, - ) - # Fallback: yield original usage chunk if adjustment fails - yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode() + ) + cost_data = await adjust_payment_for_tokens( + fresh_key, + adjustment_input, + session, + max_cost_for_model, + ) + usage_finalized = True + except Exception as e: + logger.exception( + "Error during usage finalization", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "error": str(e), + }, + ) - if not usage_finalized: - await finalize_db_only() + # Fall back so we still emit a non-zero sats cost downstream. + cost_data = { + "base_msats": 0, + "input_msats": 0, + "output_msats": 0, + "total_msats": 0, + "total_usd": 0.0, + "input_tokens": 0, + "output_tokens": 0, + } + + if usage_chunk_data is None: + if not hasattr(self, "_current_stream_id"): + self._current_stream_id = ( + f"chatcmpl-{uuid.uuid4()}" + ) + usage_chunk_data = { + "id": self._current_stream_id, + "object": "chat.completion.chunk", + "model": last_model_seen or "unknown", + "choices": [], + "usage": { + "prompt_tokens": cost_data.get( + "input_tokens", 0 + ), + "completion_tokens": cost_data.get( + "output_tokens", 0 + ), + "total_tokens": cost_data.get( + "input_tokens", 0 + ) + + cost_data.get("output_tokens", 0), + }, + } + + try: + self.inject_cost_metadata( + usage_chunk_data, cost_data, fresh_key + ) + except Exception: + logger.exception( + "Failed to inject cost metadata into streaming chunk", + extra={ + "key_hash": key.hashed_key[:8] + "...", + }, + ) + + yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode() if done_seen: yield b"data: [DONE]\n\n" @@ -926,65 +952,108 @@ class BaseUpstreamProvider: ) yield prefix + part - # Stream finished, process usage if found - if usage_chunk_data: - async with create_session() as session: - fresh_key = await session.get(key.__class__, key.hashed_key) - if fresh_key: - try: - cost_data = await adjust_payment_for_tokens( - fresh_key, - usage_chunk_data, - session, - max_cost_for_model, - ) - remaining_balance_msats = fresh_key.balance - # Merge cost into usage chunk - if ( - "response" in usage_chunk_data - and "usage" in usage_chunk_data["response"] - ): - usage_chunk_data["response"]["usage"]["cost"] = ( - cost_data.get("total_usd", 0.0) - ) - usage_chunk_data["response"]["usage"][ - "cost_sats" - ] = cost_data.get("total_msats", 0) // 1000 - usage_chunk_data["response"]["usage"][ - "remaining_balance_msats" - ] = remaining_balance_msats - elif "usage" in usage_chunk_data: - usage_chunk_data["usage"]["cost"] = cost_data.get( - "total_usd", 0.0 - ) - usage_chunk_data["usage"]["cost_sats"] = ( - cost_data.get("total_msats", 0) // 1000 - ) - usage_chunk_data["usage"][ - "remaining_balance_msats" - ] = remaining_balance_msats - - # Keep detailed cost in metadata - usage_chunk_data["metadata"] = usage_chunk_data.get( - "metadata", {} - ) - usage_chunk_data["metadata"]["routstr"] = { - "cost": cost_data + # Always emit a cost-bearing data chunk + async with create_session() as session: + fresh_key = await session.get(key.__class__, key.hashed_key) + if fresh_key: + cost_data: dict + try: + adjustment_input = ( + usage_chunk_data + if usage_chunk_data is not None + else { + "model": last_model_seen or "unknown", + "usage": None, } - usage_chunk_data["metadata"]["routstr"]["cost"][ - "sats_cost" - ] = cost_data.get("total_msats", 0) // 1000 - usage_chunk_data["metadata"]["routstr"]["cost"][ - "remaining_balance_msats" - ] = remaining_balance_msats - yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode() - usage_finalized = True - except Exception: - # Fallback: yield original usage chunk if adjustment fails - yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode() + ) + cost_data = await adjust_payment_for_tokens( + fresh_key, + adjustment_input, + session, + max_cost_for_model, + ) + usage_finalized = True + except Exception as e: + logger.exception( + "Error during Responses API usage finalization", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "error": str(e), + }, + ) + cost_data = { + "base_msats": 0, + "input_msats": 0, + "output_msats": 0, + "total_msats": 0, + "total_usd": 0.0, + "input_tokens": 0, + "output_tokens": 0, + } - if not usage_finalized: - await finalize_db_only() + if usage_chunk_data is None: + usage_chunk_data = { + "type": "response.completed", + "response": { + "model": last_model_seen or "unknown", + "usage": { + "input_tokens": cost_data.get( + "input_tokens", 0 + ), + "output_tokens": cost_data.get( + "output_tokens", 0 + ), + "total_tokens": cost_data.get( + "input_tokens", 0 + ) + + cost_data.get("output_tokens", 0), + }, + }, + "usage": { + "input_tokens": cost_data.get( + "input_tokens", 0 + ), + "output_tokens": cost_data.get( + "output_tokens", 0 + ), + "total_tokens": cost_data.get( + "input_tokens", 0 + ) + + cost_data.get("output_tokens", 0), + }, + } + + remaining_balance_msats = fresh_key.balance + sats_cost = cost_data.get("total_msats", 0) // 1000 + + if ( + "response" in usage_chunk_data + and isinstance(usage_chunk_data["response"], dict) + and "usage" in usage_chunk_data["response"] + ): + usage_chunk_data["response"]["usage"]["cost"] = ( + cost_data.get("total_usd", 0.0) + ) + usage_chunk_data["response"]["usage"][ + "cost_sats" + ] = sats_cost + usage_chunk_data["response"]["usage"][ + "remaining_balance_msats" + ] = remaining_balance_msats + + try: + self.inject_cost_metadata( + usage_chunk_data, cost_data, fresh_key + ) + except Exception: + logger.exception( + "Failed to inject cost metadata into Responses streaming chunk", + extra={ + "key_hash": key.hashed_key[:8] + "...", + }, + ) + + yield f"data: {json.dumps(usage_chunk_data)}\n\n".encode() if done_seen: yield b"data: [DONE]\n\n" @@ -1308,7 +1377,8 @@ class BaseUpstreamProvider: ) usage_finalized = True - yield f"event: cost\ndata: {json.dumps({'cost': cost_data})}\n\n".encode() + # Emit the full combined_data as the cost + yield f"event: cost\ndata: {json.dumps(combined_data)}\n\n".encode() except Exception: pass diff --git a/tests/unit/test_stream_id_injection.py b/tests/unit/test_stream_id_injection.py index 30c8323e..e19a9d3e 100644 --- a/tests/unit/test_stream_id_injection.py +++ b/tests/unit/test_stream_id_injection.py @@ -51,7 +51,15 @@ async def test_stream_with_id_injection() -> None: base.adjust_payment_for_tokens = AsyncMock( return_value={"total_usd": 0.1, "total_msats": 100} ) - base.create_session = MagicMock() + # create_session() is used as an async context manager whose entered + # value exposes an awaitable .get(). Build a mock that behaves that + # way so the post-stream cost-chunk emission can run. + mock_session = MagicMock() + mock_session.get = AsyncMock(return_value=key) + mock_ctx = MagicMock() + mock_ctx.__aenter__ = AsyncMock(return_value=mock_session) + mock_ctx.__aexit__ = AsyncMock(return_value=None) + base.create_session = MagicMock(return_value=mock_ctx) streaming_response = await provider.handle_streaming_chat_completion( response=mock_response,