diff --git a/routstr/auth.py b/routstr/auth.py index a1cba59f..6202a69a 100644 --- a/routstr/auth.py +++ b/routstr/auth.py @@ -436,6 +436,40 @@ async def adjust_payment_for_tokens( "max_cost": cost.total_msats, }, ) + # Finalize by releasing reservation and charging max cost + finalize_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values( + reserved_balance=col(ApiKey.reserved_balance) - deducted_max_cost, + balance=col(ApiKey.balance) - cost.total_msats, + total_spent=col(ApiKey.total_spent) + cost.total_msats, + ) + ) + result = await session.exec(finalize_stmt) # type: ignore[call-overload] + await session.commit() + if result.rowcount == 0: + logger.error( + "Failed to finalize max-cost payment - insufficient reserved balance", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "deducted_max_cost": deducted_max_cost, + "current_reserved_balance": key.reserved_balance, + "total_cost": cost.total_msats, + "model": model, + }, + ) + else: + await session.refresh(key) + logger.info( + "Max cost payment finalized", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "charged_amount": cost.total_msats, + "new_balance": key.balance, + "model": model, + }, + ) return cost.dict() case CostData() as cost: @@ -459,15 +493,27 @@ async def adjust_payment_for_tokens( if cost_difference == 0: logger.debug( - "No cost adjustment needed", + "Finalizing with exact reserved cost", extra={"key_hash": key.hashed_key[:8] + "...", "model": model}, ) + finalize_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values( + reserved_balance=col(ApiKey.reserved_balance) + - deducted_max_cost, + balance=col(ApiKey.balance) - total_cost_msats, + total_spent=col(ApiKey.total_spent) + total_cost_msats, + ) + ) + await session.exec(finalize_stmt) # type: ignore[call-overload] await session.commit() + await session.refresh(key) return cost.dict() # this should never happen why do we handle this??? if cost_difference > 0: - # Need to charge more + # Need to charge more than reserved, finalize by releasing reservation and charging total logger.info( "Additional charge required for token usage", extra={ @@ -479,56 +525,41 @@ async def adjust_payment_for_tokens( }, ) - # this should never happen why do we handle this??? - if key.balance < cost_difference: - logger.warning( - "Insufficient balance for token-based pricing adjustment", + finalize_stmt = ( + update(ApiKey) + .where(col(ApiKey.hashed_key) == key.hashed_key) + .values( + reserved_balance=col(ApiKey.reserved_balance) + - deducted_max_cost, + balance=col(ApiKey.balance) - total_cost_msats, + total_spent=col(ApiKey.total_spent) + total_cost_msats, + ) + ) + result = await session.exec(finalize_stmt) # type: ignore[call-overload] + await session.commit() + + if result.rowcount: + cost.total_msats = total_cost_msats + await session.refresh(key) + + logger.info( + "Finalized payment with additional charge", extra={ "key_hash": key.hashed_key[:8] + "...", - "required": cost_difference, - "available": key.balance, - "shortfall": cost_difference - key.balance, + "charged_amount": total_cost_msats, + "new_balance": key.balance, "model": model, }, ) - await session.commit() else: - # this should never happen why do we handle this??? - charge_stmt = ( - update(ApiKey) - .where(col(ApiKey.hashed_key) == key.hashed_key) - .where(col(ApiKey.balance) >= cost_difference) - .values( - balance=col(ApiKey.balance) - cost_difference, - total_spent=col(ApiKey.total_spent) + cost_difference, - ) + logger.warning( + "Failed to finalize additional charge (concurrent operation)", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "attempted_charge": total_cost_msats, + "model": model, + }, ) - result = await session.exec(charge_stmt) # type: ignore[call-overload] - await session.commit() - - if result.rowcount: - cost.total_msats = deducted_max_cost + cost_difference - await session.refresh(key) - - logger.info( - "Additional charge applied successfully", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "charged_amount": cost_difference, - "new_balance": key.balance, - "total_cost": cost.total_msats, - "model": model, - }, - ) - else: - logger.warning( - "Failed to apply additional charge (concurrent operation)", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "attempted_charge": cost_difference, - "model": model, - }, - ) else: # Refund some of the base cost refund = abs(cost_difference) diff --git a/routstr/proxy.py b/routstr/proxy.py index 5ebe7c6f..c07d6dc2 100644 --- a/routstr/proxy.py +++ b/routstr/proxy.py @@ -42,106 +42,157 @@ async def handle_streaming_chat_completion( ) async def stream_with_cost(max_cost_for_model: int) -> AsyncGenerator[bytes, None]: - # Store all chunks to analyze - stored_chunks = [] + stored_chunks: list[bytes] = [] + usage_finalized: bool = False + last_model_seen: str | None = None - async for chunk in response.aiter_bytes(): - # Store chunk for later analysis - stored_chunks.append(chunk) + async def finalize_without_usage() -> bytes | None: + nonlocal usage_finalized + if usage_finalized: + return None + async with create_session() as new_session: + fresh_key = await new_session.get(key.__class__, key.hashed_key) + if not fresh_key: + return None + try: + fallback: dict = { + "model": last_model_seen or "unknown", + "usage": None, + } + cost_data = await adjust_payment_for_tokens( + fresh_key, fallback, new_session, max_cost_for_model + ) + usage_finalized = True + logger.info( + "Finalized streaming payment without explicit usage", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "cost_data": cost_data, + "balance_after_adjustment": fresh_key.balance, + }, + ) + return f"data: {json.dumps({'cost': cost_data})}\n\n".encode() + except Exception as cost_error: + logger.error( + "Error finalizing payment without usage", + extra={ + "error": str(cost_error), + "error_type": type(cost_error).__name__, + "key_hash": key.hashed_key[:8] + "...", + }, + ) + return None - # Pass through each chunk to client - yield chunk + try: + async for chunk in response.aiter_bytes(): + stored_chunks.append(chunk) + # Opportunistically capture model id + try: + for part in re.split(b"data: ", chunk): + if not part or part.strip() in (b"[DONE]", b""): + continue + try: + obj = json.loads(part) + if isinstance(obj, dict) and obj.get("model"): + last_model_seen = str(obj.get("model")) + except json.JSONDecodeError: + pass + except Exception: + pass - logger.debug( - "Streaming completed, analyzing usage data", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "chunks_count": len(stored_chunks), - }, - ) + yield chunk - # Process stored chunks to find usage data - # Start from the end and work backwards - for i in range(len(stored_chunks) - 1, -1, -1): - chunk = stored_chunks[i] - if not chunk or chunk == b"": - continue + logger.debug( + "Streaming completed, analyzing usage data", + extra={ + "key_hash": key.hashed_key[:8] + "...", + "chunks_count": len(stored_chunks), + }, + ) - try: - # Split by "data: " to get individual SSE events - events = re.split(b"data: ", chunk) - for event_data in events: - if ( - not event_data - or event_data.strip() == b"[DONE]" - or event_data.strip() == b"" - ): - continue + # Process stored chunks to find usage data from the tail + for i in range(len(stored_chunks) - 1, -1, -1): + chunk = stored_chunks[i] + if not chunk: + continue + try: + events = re.split(b"data: ", chunk) + for event_data in events: + if not event_data or event_data.strip() in (b"[DONE]", b""): + continue + try: + data = json.loads(event_data) + if isinstance(data, dict) and data.get("model"): + last_model_seen = str(data.get("model")) + if isinstance(data, dict) and isinstance( + data.get("usage"), dict + ): + async with create_session() as new_session: + fresh_key = await new_session.get( + key.__class__, key.hashed_key + ) + if fresh_key: + try: + cost_data = await adjust_payment_for_tokens( + fresh_key, + data, + new_session, + max_cost_for_model, + ) + usage_finalized = True + logger.info( + "Token adjustment completed for streaming", + extra={ + "key_hash": key.hashed_key[:8] + + "...", + "cost_data": cost_data, + "balance_after_adjustment": fresh_key.balance, + }, + ) + yield f"data: {json.dumps({'cost': cost_data})}\n\n".encode() + except Exception as cost_error: + logger.error( + "Error adjusting payment for streaming tokens", + extra={ + "error": str(cost_error), + "error_type": type( + cost_error + ).__name__, + "key_hash": key.hashed_key[:8] + + "...", + }, + ) + break + except json.JSONDecodeError: + continue + except Exception as e: + logger.error( + "Error processing streaming response chunk", + extra={ + "error": str(e), + "error_type": type(e).__name__, + "key_hash": key.hashed_key[:8] + "...", + }, + ) - try: - data = json.loads(event_data) - if ( - "usage" in data - and data["usage"] is not None - and isinstance(data["usage"], dict) - ): - logger.info( - "Found usage data in streaming response", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "usage_data": data["usage"], - "model": data.get("model", "unknown"), - }, - ) + # If we reach here without finding usage, finalize with max-cost + if not usage_finalized: + maybe_cost_event = await finalize_without_usage() + if maybe_cost_event is not None: + yield maybe_cost_event - # Found usage data, calculate cost - # Create a new session for this operation - async with create_session() as new_session: - # Re-fetch the key in the new session - fresh_key = await new_session.get( - key.__class__, key.hashed_key - ) - if fresh_key: - try: - cost_data = await adjust_payment_for_tokens( - fresh_key, - data, - new_session, - max_cost_for_model, - ) - logger.info( - "Token adjustment completed for streaming", - extra={ - "key_hash": key.hashed_key[:8] + "...", - "cost_data": cost_data, - "balance_after_adjustment": fresh_key.balance, - }, - ) - # Format as SSE and yield - cost_json = json.dumps({"cost": cost_data}) - yield f"data: {cost_json}\n\n".encode() - except Exception as cost_error: - logger.error( - "Error adjusting payment for streaming tokens", - extra={ - "error": str(cost_error), - "error_type": type(cost_error).__name__, - "key_hash": key.hashed_key[:8] + "...", - }, - ) - break - except json.JSONDecodeError: - continue - - except Exception as e: - logger.error( - "Error processing streaming response chunk", - extra={ - "error": str(e), - "error_type": type(e).__name__, - "key_hash": key.hashed_key[:8] + "...", - }, - ) + except Exception as stream_error: + # On stream interruption, still finalize reservation with max-cost + logger.warning( + "Streaming interrupted; finalizing without usage", + extra={ + "error": str(stream_error), + "error_type": type(stream_error).__name__, + "key_hash": key.hashed_key[:8] + "...", + }, + ) + await finalize_without_usage() + raise return StreamingResponse( stream_with_cost(max_cost_for_model),