mirror of
https://github.com/Routstr/routstr-core.git
synced 2026-07-31 15:56:14 +00:00
Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9f55da9bb8 | ||
|
|
16dce9ea81 | ||
|
|
963ee04619 | ||
|
|
b874b1f01c | ||
|
|
b6cca3d3a0 | ||
|
|
86ebc84f4c |
+30
-13
@@ -5,6 +5,7 @@ import hashlib
|
||||
import json
|
||||
import re
|
||||
import traceback
|
||||
import uuid
|
||||
from collections.abc import AsyncGenerator
|
||||
from typing import Mapping
|
||||
|
||||
@@ -517,20 +518,32 @@ class BaseUpstreamProvider:
|
||||
continue
|
||||
|
||||
try:
|
||||
obj = json.loads(part)
|
||||
if isinstance(obj, dict):
|
||||
if obj.get("model"):
|
||||
last_model_seen = str(obj.get("model"))
|
||||
if requested_model:
|
||||
obj["model"] = requested_model
|
||||
|
||||
if isinstance(obj.get("usage"), dict):
|
||||
# Hold this chunk back to merge cost later
|
||||
usage_chunk_data = obj
|
||||
# Only parse if it looks like a JSON object to avoid SSE control messages or partials
|
||||
if part.strip().startswith(b"{") and part.strip().endswith(
|
||||
b"}"
|
||||
):
|
||||
obj = json.loads(part)
|
||||
if isinstance(obj, dict):
|
||||
if obj.get("model"):
|
||||
last_model_seen = str(obj.get("model"))
|
||||
if requested_model:
|
||||
obj["model"] = requested_model
|
||||
if (
|
||||
"id" not in obj
|
||||
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 isinstance(obj.get("usage"), dict):
|
||||
usage_chunk_data = obj
|
||||
continue
|
||||
yield b"data: " + json.dumps(obj).encode() + b"\n\n"
|
||||
continue
|
||||
yield b"data: " + json.dumps(obj).encode() + b"\n\n"
|
||||
continue
|
||||
except json.JSONDecodeError:
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
prefix = (
|
||||
@@ -663,6 +676,8 @@ 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()}"
|
||||
|
||||
cost_data = await adjust_payment_for_tokens(
|
||||
key, response_json, session, deducted_max_cost
|
||||
@@ -992,6 +1007,8 @@ 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()}"
|
||||
|
||||
cost_data = await adjust_payment_for_tokens(
|
||||
key, response_json, session, deducted_max_cost
|
||||
|
||||
@@ -155,6 +155,20 @@ async def swap_to_primary_mint(
|
||||
raise ValueError("Invalid unit")
|
||||
primary_wallet = await get_wallet(settings.primary_mint, settings.primary_mint_unit)
|
||||
|
||||
# If the token is already from the primary mint, we don't need to swap
|
||||
# and we definitely don't want to calculate or pay fees.
|
||||
if token_obj.mint == settings.primary_mint:
|
||||
logger.info(
|
||||
"swap_to_primary_mint: token already on primary mint, skipping swap",
|
||||
extra={
|
||||
"mint": token_obj.mint,
|
||||
"amount": token_amount,
|
||||
"unit": token_obj.unit,
|
||||
},
|
||||
)
|
||||
await token_wallet.split(proofs=token_obj.proofs, amount=0, include_fees=True)
|
||||
return token_amount, token_obj.unit, token_obj.mint
|
||||
|
||||
minted_amount = await _calculate_swap_amount(
|
||||
amount_msat,
|
||||
token_obj.unit,
|
||||
|
||||
@@ -0,0 +1,109 @@
|
||||
import json
|
||||
from collections.abc import AsyncGenerator
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from routstr.core.db import ApiKey
|
||||
from routstr.upstream.base import BaseUpstreamProvider
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_with_id_injection() -> None:
|
||||
"""Test that stream_with_cost correctly injects IDs into complete JSON chunks but skips partials."""
|
||||
provider = BaseUpstreamProvider(
|
||||
base_url="https://api.example.com", api_key="test_key"
|
||||
)
|
||||
|
||||
# Mock response with mixed chunks:
|
||||
# 1. Complete JSON without ID
|
||||
# 2. Partial JSON (should be passed through)
|
||||
# 3. Complete JSON with ID (should be preserved or updated if requested_model is set)
|
||||
# 4. [DONE] message
|
||||
chunks = [
|
||||
b'data: {"choices": [{"delta": {"content": "Hello"}}]}\n\n',
|
||||
b'data: {"choices": [{"delta": {"content": "', # Partial
|
||||
b'world"}}]}\n\n',
|
||||
b'data: {"id": "existing-id", "choices": [{"delta": {"content": "!"}}]}\n\n',
|
||||
b"data: [DONE]\n\n",
|
||||
]
|
||||
|
||||
async def aiter_bytes() -> AsyncGenerator[bytes, None]:
|
||||
for chunk in chunks:
|
||||
yield chunk
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.headers = {"content-type": "text/event-stream"}
|
||||
mock_response.aiter_bytes = aiter_bytes
|
||||
|
||||
key = MagicMock(spec=ApiKey)
|
||||
key.hashed_key = "test_hash"
|
||||
key.balance = 1000
|
||||
|
||||
background_tasks = MagicMock()
|
||||
|
||||
# We need to mock adjust_payment_for_tokens since it's called at the end
|
||||
with MagicMock():
|
||||
from routstr.upstream import base
|
||||
|
||||
# Mocking the module-level function used in the generator
|
||||
base.adjust_payment_for_tokens = AsyncMock(
|
||||
return_value={"total_usd": 0.1, "total_msats": 100}
|
||||
)
|
||||
base.create_session = MagicMock()
|
||||
|
||||
streaming_response = await provider.handle_streaming_chat_completion(
|
||||
response=mock_response,
|
||||
key=key,
|
||||
max_cost_for_model=100,
|
||||
background_tasks=background_tasks,
|
||||
requested_model="test-model",
|
||||
)
|
||||
|
||||
results = []
|
||||
async for chunk in streaming_response.body_iterator:
|
||||
results.append(chunk)
|
||||
|
||||
# Parse results
|
||||
parsed_results = []
|
||||
for r in results:
|
||||
if isinstance(r, bytes) and r.startswith(b"data: "):
|
||||
data = r[6:].decode().strip()
|
||||
if data == "[DONE]":
|
||||
parsed_results.append(data)
|
||||
else:
|
||||
try:
|
||||
parsed_results.append(json.loads(data))
|
||||
except (json.JSONDecodeError, UnicodeDecodeError):
|
||||
parsed_results.append(
|
||||
data
|
||||
) # Keep as string if it failed to parse
|
||||
|
||||
# Verifications
|
||||
# 1. First chunk should have an injected ID and the requested model
|
||||
assert isinstance(parsed_results[0], dict)
|
||||
assert "id" in parsed_results[0]
|
||||
assert parsed_results[0]["id"].startswith("chatcmpl-")
|
||||
assert parsed_results[0]["model"] == "test-model"
|
||||
|
||||
# 2. Second chunk was partial, should be passed as-is
|
||||
# In current implementation, re.split(b"data: ", b'data: {...') gives ['', '{...']
|
||||
# The first empty part is skipped. The second part is processed.
|
||||
|
||||
# Check that we have results
|
||||
assert len(parsed_results) >= 4
|
||||
|
||||
# Find the chunk that was "existing-id"
|
||||
id_chunk = next(
|
||||
r
|
||||
for r in parsed_results
|
||||
if isinstance(r, dict)
|
||||
and "choices" in r
|
||||
and r["choices"][0]["delta"].get("content") == "!"
|
||||
)
|
||||
assert id_chunk["id"] == parsed_results[0]["id"]
|
||||
assert id_chunk["model"] == "test-model"
|
||||
|
||||
# 4. [DONE] should be there
|
||||
assert "[DONE]" in parsed_results
|
||||
@@ -220,6 +220,33 @@ async def test_recieve_token_untrusted_mint() -> None:
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.asyncio
|
||||
async def test_swap_to_primary_mint_already_on_primary() -> None:
|
||||
from routstr.core.settings import settings
|
||||
from routstr.wallet import swap_to_primary_mint
|
||||
|
||||
mock_token = Mock()
|
||||
mock_token.mint = settings.primary_mint
|
||||
mock_token.amount = 1000
|
||||
mock_token.unit = "sat"
|
||||
mock_token.proofs = []
|
||||
|
||||
mock_token_wallet = Mock()
|
||||
mock_token_wallet.split = AsyncMock(return_value=None)
|
||||
mock_token_wallet.request_mint = AsyncMock()
|
||||
mock_token_wallet.melt_quote = AsyncMock()
|
||||
|
||||
with patch("routstr.wallet.get_wallet", AsyncMock(return_value=mock_token_wallet)):
|
||||
amount, unit, mint = await swap_to_primary_mint(mock_token, mock_token_wallet)
|
||||
|
||||
assert amount == 1000
|
||||
assert unit == "sat"
|
||||
assert mint == settings.primary_mint
|
||||
mock_token_wallet.split.assert_called_once()
|
||||
mock_token_wallet.request_mint.assert_not_called()
|
||||
mock_token_wallet.melt_quote.assert_not_called()
|
||||
|
||||
|
||||
async def test_swap_to_primary_mint_success() -> None:
|
||||
"""Test successful swap with dynamic fee calculation."""
|
||||
from routstr.wallet import swap_to_primary_mint
|
||||
|
||||
Reference in New Issue
Block a user