diff --git a/routstr/upstream/base.py b/routstr/upstream/base.py index 8f46d90b..2f65ebd3 100644 --- a/routstr/upstream/base.py +++ b/routstr/upstream/base.py @@ -2612,6 +2612,11 @@ class BaseUpstreamProvider: ) -> dict: return await messages_dispatch.aggregate_anthropic_events_to_message(iterator) + def transform_messages_stream( + self, stream: AsyncIterator[Any] + ) -> AsyncIterator[Any]: + return stream + def adapt_messages_request(self, body: dict, model_obj: Model) -> str: """Rewrite an allowlisted /v1/messages body for this upstream. @@ -2637,6 +2642,7 @@ class BaseUpstreamProvider: provider_prefix=self.get_litellm_provider_prefix(), transform_model_name=self.transform_model_name, adapt_request=lambda body: self.adapt_messages_request(body, model_obj), + transform_stream=self.transform_messages_stream, log_extra=log_extra, ) diff --git a/routstr/upstream/messages_dispatch.py b/routstr/upstream/messages_dispatch.py index 6de5e57f..d9df0277 100644 --- a/routstr/upstream/messages_dispatch.py +++ b/routstr/upstream/messages_dispatch.py @@ -406,16 +406,9 @@ def annotate_event(event: dict, requested_model: str | None) -> AnnotatedEvent: _coerce_float(root_cost_details.get("output_cost")), ) - event_type = str(event.get("type") or "") - payload = json.dumps(event) - if event_type: - sse_bytes = f"event: {event_type}\ndata: {payload}\n\n".encode() - else: - sse_bytes = f"data: {payload}\n\n".encode() - return AnnotatedEvent( event, - sse_bytes, + encode_sse(event), in_tokens, out_tokens, cache_read_tokens, @@ -427,6 +420,14 @@ def annotate_event(event: dict, requested_model: str | None) -> AnnotatedEvent: ) +def encode_sse(event: dict) -> bytes: + event_type = str(event.get("type") or "") + payload = json.dumps(event) + if event_type: + return f"event: {event_type}\ndata: {payload}\n\n".encode() + return f"data: {payload}\n\n".encode() + + async def stream_annotated_events( iterator: AsyncIterator[Any], requested_model: str | None, @@ -493,6 +494,7 @@ async def dispatch_anthropic_messages( provider_prefix: str, transform_model_name: Callable[[str], str], adapt_request: Callable[[dict], str] | None = None, + transform_stream: Callable[[AsyncIterator[Any]], AsyncIterator[Any]] | None = None, log_extra: dict[str, Any] | None = None, ) -> tuple[bool, Any, str | None]: """Call ``litellm.anthropic.messages.acreate`` and return @@ -505,6 +507,10 @@ async def dispatch_anthropic_messages( may rewrite it in place and returns a suffix for the upstream model name, which is how a provider expresses a feature litellm would otherwise translate into a parameter the upstream rejects. + + ``transform_stream`` rewrites the upstream event stream before it is + aggregated or handed to the client, so a provider can repair events + litellm translates faithfully but clients cannot use. """ if not request_body: raise UpstreamError("Missing request body for /v1/messages", status_code=400) @@ -640,6 +646,9 @@ async def dispatch_anthropic_messages( from_upstream_response=True, ) from exc + if transform_stream is not None and hasattr(result, "__aiter__"): + result = transform_stream(cast(AsyncIterator[Any], result)) + if not client_stream and hasattr(result, "__aiter__"): # Client asked for a non-streaming response but we always stream # from upstream — drain the events into a single Anthropic Message diff --git a/routstr/upstream/venice.py b/routstr/upstream/venice.py index 379a4a66..15a476fa 100644 --- a/routstr/upstream/venice.py +++ b/routstr/upstream/venice.py @@ -1,13 +1,16 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any +from collections.abc import AsyncGenerator, AsyncIterator +from typing import TYPE_CHECKING, Any, cast import httpx from ..core.exceptions import UpstreamError from ..core.logging import get_logger from ..payment.models import Architecture, Model, Pricing, TopProvider +from . import messages_dispatch from .base import BaseUpstreamProvider +from .stream_ownership import aclose_if_needed if TYPE_CHECKING: from ..core.db import UpstreamProviderRow @@ -50,6 +53,73 @@ _UNENFORCEABLE_WEB_SEARCH_KEYS = frozenset( {"allowed_domains", "blocked_domains", "user_location"} ) +# Venice streams OpenAI reasoning models' encrypted reasoning as a trailing +# ``reasoning_content`` delta carrying this marker. litellm turns it into a +# plaintext ``thinking`` block after the answer, which clients render as +# gibberish and which makes Claude Code report an empty final result. +_ENCRYPTED_REASONING_MARKER = "__ENCRYPTED_REASONING__" + + +async def _drop_encrypted_reasoning( + upstream: AsyncIterator[Any], +) -> AsyncGenerator[bytes, None]: + """A thinking block's start carries no text, so it is held until its first + delta shows whether it is the encrypted payload; later indices shift down + to close the gap.""" + encode = messages_dispatch.encode_sse + sse_buffer = b"" + dropped: set[int] = set() + held: list[dict] | None = None + held_index: int | None = None + + def shift(event: dict) -> dict: + index = event.get("index") + if not isinstance(index, int): + return event + gap = sum(1 for d in dropped if d < index) + return {**event, "index": index - gap} if gap else event + + try: + async for chunk in upstream: + events, sse_buffer = messages_dispatch.events_from_chunk(chunk, sse_buffer) + for event in events: + etype = event.get("type") + index = event.get("index") + if held is not None: + delta = event.get("delta") or {} + is_own_delta = ( + index == held_index and etype == "content_block_delta" + ) + thinking = str(delta.get("thinking") or "") + if is_own_delta and thinking.startswith( + _ENCRYPTED_REASONING_MARKER + ): + dropped.add(cast(int, index)) + held = None + continue + if is_own_delta and not thinking: + held.append(event) + continue + for pending in held: + yield encode(shift(pending)) + held = None + if index in dropped: + continue + block = event.get("content_block") or {} + if ( + etype == "content_block_start" + and block.get("type") == "thinking" + and not block.get("thinking") + ): + held, held_index = [event], index + continue + yield encode(shift(event)) + if held is not None: + for pending in held: + yield encode(shift(pending)) + finally: + await aclose_if_needed(upstream) + def _is_web_search_tool(tool: Any) -> bool: """An Anthropic server-side web-search tool, by either of its markers. @@ -66,6 +136,35 @@ def _is_web_search_tool(tool: Any) -> bool: ) or tool.get("name") == "web_search" +def _merge_cache_marked_system(body: dict) -> None: + """Venice rejects an OpenAI ``system`` message with two or more text parts + when any part carries ``cache_control`` (``400 system: text content blocks + must contain non-whitespace text``), even though every part is non-blank. + Claude Code always sends that shape. A single marked block is accepted and + still caches, so the prefix stays cacheable under the last marker. + """ + system = body.get("system") + if not isinstance(system, list) or len(system) < 2: + return + if not all( + isinstance(block, dict) + and block.get("type") == "text" + and isinstance(block.get("text"), str) + for block in system + ): + return + markers = [block["cache_control"] for block in system if block.get("cache_control")] + if not markers: + return + body["system"] = [ + { + "type": "text", + "text": "\n\n".join(block["text"] for block in system), + "cache_control": markers[-1], + } + ] + + def _usd(entry: Any) -> float | None: """Read the USD leg of a Venice ``{usd, diem}`` price pair.""" if isinstance(entry, dict): @@ -114,7 +213,16 @@ class VeniceUpstreamProvider(BaseUpstreamProvider): def transform_model_name(self, model_id: str) -> str: return model_id.removeprefix("venice/") + def transform_messages_stream( + self, stream: AsyncIterator[Any] + ) -> AsyncIterator[Any]: + return _drop_encrypted_reasoning(stream) + def adapt_messages_request(self, body: dict, model_obj: Model) -> str: + _merge_cache_marked_system(body) + return self._adapt_web_search(body) + + def _adapt_web_search(self, body: dict) -> str: """Trade an Anthropic web-search tool for Venice's own search switch. Left in the body, litellm's Anthropic adapter rewrites the tool into a diff --git a/tests/unit/test_venice_encrypted_reasoning.py b/tests/unit/test_venice_encrypted_reasoning.py new file mode 100644 index 00000000..43af7614 --- /dev/null +++ b/tests/unit/test_venice_encrypted_reasoning.py @@ -0,0 +1,210 @@ +import json +from collections.abc import AsyncIterator +from typing import Any +from unittest.mock import AsyncMock, patch + +import pytest + +from routstr.upstream import messages_dispatch +from routstr.upstream.base import BaseUpstreamProvider +from routstr.upstream.venice import VeniceUpstreamProvider, _drop_encrypted_reasoning + +from .test_venice_web_search import _model + +ENCRYPTED = "__ENCRYPTED_REASONING__id=rs_0b04\ngAAAAABqvDJD" + + +def _block(index: int, block: dict, deltas: list[dict]) -> list[dict]: + return [ + {"type": "content_block_start", "index": index, "content_block": block}, + *({"type": "content_block_delta", "index": index, "delta": d} for d in deltas), + {"type": "content_block_stop", "index": index}, + ] + + +def _thinking(index: int, text: str) -> list[dict]: + return _block( + index, + {"type": "thinking", "thinking": "", "signature": ""}, + [{"type": "thinking_delta", "thinking": text}], + ) + + +def _text(index: int, text: str) -> list[dict]: + return _block( + index, + {"type": "text", "text": ""}, + [{"type": "text_delta", "text": text}], + ) + + +def _tool(index: int) -> list[dict]: + return _block( + index, + {"type": "tool_use", "id": "call_1", "name": "Bash", "input": {}}, + [{"type": "input_json_delta", "partial_json": '{"command":"ls"}'}], + ) + + +def _message(blocks: list[dict], stop_reason: str = "end_turn") -> list[dict]: + return [ + { + "type": "message_start", + "message": {"id": "msg_1", "role": "assistant", "content": []}, + }, + *blocks, + {"type": "message_delta", "delta": {"stop_reason": stop_reason}}, + {"type": "message_stop"}, + ] + + +async def _upstream(events: list[dict], *, split: bool = False) -> AsyncIterator[Any]: + payload = b"".join(messages_dispatch.encode_sse(e) for e in events) + if split: + for i in range(0, len(payload), 7): + yield payload[i : i + 7] + else: + yield payload + + +async def _filtered(events: list[dict], **kwargs: Any) -> list[dict]: + buffer = b"" + out: list[dict] = [] + async for chunk in _drop_encrypted_reasoning(_upstream(events, **kwargs)): + parsed, buffer = messages_dispatch.events_from_chunk(chunk, buffer) + out.extend(parsed) + return out + + +def _starts(events: list[dict]) -> list[tuple[int, str]]: + return [ + (e["index"], e["content_block"]["type"]) + for e in events + if e["type"] == "content_block_start" + ] + + +@pytest.mark.asyncio +async def test_trailing_encrypted_reasoning_is_dropped() -> None: + events = _message([*_text(0, "a.txt contains: hello"), *_thinking(1, ENCRYPTED)]) + + out = await _filtered(events) + + assert _starts(out) == [(0, "text")] + assert all(ENCRYPTED not in json.dumps(e) for e in out) + assert out[-2]["delta"]["stop_reason"] == "end_turn" + + +@pytest.mark.asyncio +async def test_leading_encrypted_reasoning_closes_index_gap() -> None: + events = _message( + [*_thinking(0, ENCRYPTED), *_text(1, "hi"), *_tool(2)], "tool_use" + ) + + out = await _filtered(events, split=True) + + assert _starts(out) == [(0, "text"), (1, "tool_use")] + assert {e["index"] for e in out if "index" in e} == {0, 1} + + +@pytest.mark.asyncio +async def test_plaintext_thinking_is_kept_in_order() -> None: + events = _message([*_thinking(0, "Let me list files."), *_tool(1)], "tool_use") + + out = await _filtered(events) + + assert out == events + + +@pytest.mark.asyncio +async def test_thinking_start_without_delta_is_flushed() -> None: + events = _message( + [ + { + "type": "content_block_start", + "index": 0, + "content_block": {"type": "thinking", "thinking": "", "signature": ""}, + }, + {"type": "content_block_stop", "index": 0}, + *_text(1, "ok"), + ] + ) + + out = await _filtered(events) + + assert out == events + + +@pytest.mark.asyncio +async def test_aggregated_message_ends_with_answer_text() -> None: + events = _message([*_text(0, "hello"), *_thinking(1, ENCRYPTED)]) + + message = await messages_dispatch.aggregate_anthropic_events_to_message( + _drop_encrypted_reasoning(_upstream(events)) + ) + + assert [b["type"] for b in message["content"]] == ["text"] + assert message["content"][0]["text"] == "hello" + + +async def _dispatched_blocks( + provider: BaseUpstreamProvider, *, stream: bool +) -> list[str]: + events = _message([*_text(0, "hello"), *_thinking(1, ENCRYPTED)]) + with patch( + "litellm.anthropic.messages.acreate", + new=AsyncMock(return_value=_upstream(events)), + ): + _, result, _ = await provider._dispatch_anthropic_messages( + request_body=json.dumps( + { + "model": "x", + "stream": stream, + "max_tokens": 64, + "messages": [{"role": "user", "content": "hi"}], + } + ).encode(), + model_obj=_model(), + ) + if not stream: + return [b["type"] for b in result["content"]] + buffer = b"" + out: list[dict] = [] + async for chunk in result: + parsed, buffer = messages_dispatch.events_from_chunk(chunk, buffer) + out.extend(parsed) + return [t for _, t in _starts(out)] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stream", [True, False]) +async def test_venice_dispatch_drops_encrypted_reasoning(stream: bool) -> None: + provider = VeniceUpstreamProvider(api_key="sk-test") + + assert await _dispatched_blocks(provider, stream=stream) == ["text"] + + +@pytest.mark.asyncio +async def test_other_providers_keep_thinking_blocks() -> None: + provider = BaseUpstreamProvider(base_url="https://example.com/v1", api_key="k") + + assert await _dispatched_blocks(provider, stream=True) == ["text", "thinking"] + + +@pytest.mark.asyncio +async def test_closing_the_filter_closes_upstream() -> None: + closed = False + + async def upstream() -> AsyncIterator[bytes]: + nonlocal closed + try: + for event in _message(_text(0, "hello")): + yield messages_dispatch.encode_sse(event) + finally: + closed = True + + filtered = _drop_encrypted_reasoning(upstream()) + await filtered.__anext__() + await filtered.aclose() + + assert closed diff --git a/tests/unit/test_venice_system_cache.py b/tests/unit/test_venice_system_cache.py new file mode 100644 index 00000000..6b8cc67f --- /dev/null +++ b/tests/unit/test_venice_system_cache.py @@ -0,0 +1,86 @@ +from __future__ import annotations + +import pytest + +from routstr.upstream.venice import VeniceUpstreamProvider + +from .test_venice_web_search import _body, _dispatch + +EPHEMERAL = {"type": "ephemeral"} + +CLAUDE_CODE_SYSTEM = [ + { + "type": "text", + "text": "x-anthropic-billing-header: cc_version=2.1.281; cc_entrypoint=cli;", + }, + {"type": "text", "text": "You are a Claude agent.", "cache_control": EPHEMERAL}, + { + "type": "text", + "text": "\nYou are an interactive agent.", + "cache_control": {"type": "ephemeral", "ttl": "1h"}, + }, +] + + +@pytest.mark.asyncio +async def test_cache_marked_multi_block_system_is_merged_into_one_block() -> None: + provider = VeniceUpstreamProvider(api_key="sk-test") + + kwargs = await _dispatch(provider, _body(system=CLAUDE_CODE_SYSTEM)) + + assert kwargs["system"] == [ + { + "type": "text", + "text": ( + "x-anthropic-billing-header: cc_version=2.1.281; cc_entrypoint=cli;" + "\n\nYou are a Claude agent.\n\n\nYou are an interactive agent." + ), + "cache_control": {"type": "ephemeral", "ttl": "1h"}, + } + ] + + +@pytest.mark.asyncio +async def test_unmarked_multi_block_system_is_untouched() -> None: + provider = VeniceUpstreamProvider(api_key="sk-test") + system = [{"type": "text", "text": "A."}, {"type": "text", "text": "B."}] + + kwargs = await _dispatch(provider, _body(system=system)) + + assert kwargs["system"] == system + + +@pytest.mark.asyncio +async def test_single_marked_block_and_string_system_are_untouched() -> None: + provider = VeniceUpstreamProvider(api_key="sk-test") + single = [{"type": "text", "text": "A.", "cache_control": EPHEMERAL}] + + assert (await _dispatch(provider, _body(system=single)))["system"] == single + assert (await _dispatch(provider, _body(system="A.")))["system"] == "A." + + +@pytest.mark.asyncio +async def test_message_and_tool_cache_markers_are_kept() -> None: + provider = VeniceUpstreamProvider(api_key="sk-test") + messages = [ + { + "role": "user", + "content": [{"type": "text", "text": "hi", "cache_control": EPHEMERAL}], + } + ] + tools = [ + { + "name": "Bash", + "description": "Run a command", + "input_schema": {"type": "object", "properties": {}}, + "cache_control": EPHEMERAL, + } + ] + + kwargs = await _dispatch( + provider, + _body(system=CLAUDE_CODE_SYSTEM, messages=messages, tools=tools), + ) + + assert kwargs["messages"] == messages + assert kwargs["tools"] == tools