From efa9f36c3ad87b35b7a28aa5cd3138069c0a8f24 Mon Sep 17 00:00:00 2001 From: Alishahryar1 Date: Fri, 24 Apr 2026 17:27:26 -0700 Subject: [PATCH] Revert "Refactor native Anthropic messages providers (#147)" This reverts commit ffa823722052ac4d914b36cbbf70eaa8cb29bc66. --- providers/anthropic_compat.py | 279 ------------ providers/llamacpp/client.py | 156 ++++++- providers/lmstudio/client.py | 156 ++++++- providers/open_router/client.py | 169 +------ providers/open_router/request.py | 98 +--- tests/providers/test_anthropic_compat.py | 203 --------- tests/providers/test_llamacpp.py | 9 +- tests/providers/test_lmstudio.py | 9 +- tests/providers/test_open_router.py | 542 ++++++++++++----------- 9 files changed, 617 insertions(+), 1004 deletions(-) delete mode 100644 providers/anthropic_compat.py delete mode 100644 tests/providers/test_anthropic_compat.py diff --git a/providers/anthropic_compat.py b/providers/anthropic_compat.py deleted file mode 100644 index 68e7fa8c..00000000 --- a/providers/anthropic_compat.py +++ /dev/null @@ -1,279 +0,0 @@ -"""Shared base for providers with native Anthropic Messages endpoints.""" - -import json -from collections.abc import AsyncIterator, Iterator -from typing import Any, Literal - -import httpx -from loguru import logger - -from providers.base import BaseProvider, ProviderConfig -from providers.common import get_user_facing_error_message, map_error -from providers.rate_limit import GlobalRateLimiter - -ANTHROPIC_DEFAULT_MAX_TOKENS = 81920 -StreamChunkMode = Literal["line", "event"] - - -class AnthropicMessagesProvider(BaseProvider): - """Base class for providers that stream from an Anthropic-compatible endpoint.""" - - stream_chunk_mode: StreamChunkMode = "line" - - def __init__( - self, - config: ProviderConfig, - *, - provider_name: str, - default_base_url: str, - ): - super().__init__(config) - self._provider_name = provider_name - self._api_key = config.api_key - self._base_url = (config.base_url or default_base_url).rstrip("/") - self._global_rate_limiter = GlobalRateLimiter.get_instance( - rate_limit=config.rate_limit, - rate_window=config.rate_window, - max_concurrency=config.max_concurrency, - ) - self._client = httpx.AsyncClient( - base_url=self._base_url, - proxy=config.proxy or None, - timeout=httpx.Timeout( - config.http_read_timeout, - connect=config.http_connect_timeout, - read=config.http_read_timeout, - write=config.http_write_timeout, - ), - ) - - async def cleanup(self) -> None: - """Release HTTP client resources.""" - await self._client.aclose() - - def _request_headers(self) -> dict[str, str]: - """Return headers for the native messages request.""" - return {"Content-Type": "application/json"} - - def _build_request_body(self, request: Any) -> dict: - """Build a native Anthropic request body.""" - thinking_enabled = self._is_thinking_enabled(request) - body = request.model_dump(exclude_none=True) - - body.pop("extra_body", None) - body.pop("original_model", None) - body.pop("resolved_provider_model", None) - - if "thinking" in body: - thinking_cfg = body.pop("thinking") - if ( - thinking_enabled - and isinstance(thinking_cfg, dict) - and thinking_cfg.get("enabled") - ): - body["thinking"] = {"type": "enabled"} - - if "max_tokens" not in body: - body["max_tokens"] = ANTHROPIC_DEFAULT_MAX_TOKENS - - return body - - async def _send_stream_request(self, body: dict) -> httpx.Response: - """Create a streaming messages response.""" - request = self._client.build_request( - "POST", - "/messages", - json=body, - headers=self._request_headers(), - ) - return await self._client.send(request, stream=True) - - async def _raise_for_status( - self, response: httpx.Response, *, req_tag: str - ) -> None: - """Raise for non-200 responses after logging the upstream body if available.""" - try: - response.raise_for_status() - except httpx.HTTPStatusError as error: - response_text = await self._read_error_body(response) - if response_text: - logger.error( - "{}_ERROR:{} HTTP {}: {}", - self._provider_name, - req_tag, - response.status_code, - response_text, - ) - raise error - - async def _read_error_body(self, response: httpx.Response) -> str: - """Read a response body for diagnostics when the test double supports it.""" - aread = getattr(response, "aread", None) - if aread is None: - return "" - body = await aread() - if isinstance(body, bytes): - return body.decode("utf-8", errors="replace") - return str(body) - - async def _iter_sse_lines(self, response: httpx.Response) -> AsyncIterator[str]: - """Yield raw SSE line chunks preserving existing local provider behavior.""" - async for line in response.aiter_lines(): - if line: - yield f"{line}\n" - else: - yield "\n" - - async def _iter_sse_events(self, response: httpx.Response) -> AsyncIterator[str]: - """Group line-delimited SSE responses into full SSE events.""" - event_lines: list[str] = [] - async for line in response.aiter_lines(): - if line: - event_lines.append(line) - continue - if event_lines: - yield "\n".join(event_lines) + "\n\n" - event_lines.clear() - if event_lines: - yield "\n".join(event_lines) + "\n\n" - - def _new_stream_state(self, request: Any, *, thinking_enabled: bool) -> Any: - """Return per-stream provider state for event transformation.""" - return None - - def _transform_stream_event( - self, - event: str, - state: Any, - *, - thinking_enabled: bool, - ) -> str | None: - """Transform or drop a grouped SSE event before yielding it downstream.""" - return event - - def _format_error_message(self, base_message: str, request_id: str | None) -> str: - """Apply provider-specific request id formatting to an error message.""" - if request_id: - return f"{base_message}\nRequest ID: {request_id}" - return base_message - - def _get_error_message(self, error: Exception, request_id: str | None) -> str: - """Map an exception into a user-facing provider error message.""" - mapped_error = map_error(error) - if getattr(mapped_error, "status_code", None) == 405: - base_message = ( - f"Upstream provider {self._provider_name} rejected the request method " - "or endpoint (HTTP 405)." - ) - else: - base_message = get_user_facing_error_message( - mapped_error, read_timeout_s=self._config.http_read_timeout - ) - return self._format_error_message(base_message, request_id) - - def _emit_error_events( - self, - *, - request: Any, - input_tokens: int, - error_message: str, - sent_any_event: bool, - ) -> Iterator[str]: - """Emit a native Anthropic error event.""" - error_event = { - "type": "error", - "error": {"type": "api_error", "message": error_message}, - } - yield f"event: error\ndata: {json.dumps(error_event)}\n\n" - - async def _iter_stream_chunks( - self, - response: httpx.Response, - *, - state: Any, - thinking_enabled: bool, - ) -> AsyncIterator[str]: - """Yield stream chunks according to the provider's observable chunk shape.""" - if self.stream_chunk_mode == "line": - async for chunk in self._iter_sse_lines(response): - yield chunk - return - - async for event in self._iter_sse_events(response): - output_event = self._transform_stream_event( - event, - state, - thinking_enabled=thinking_enabled, - ) - if output_event is not None: - yield output_event - - async def stream_response( - self, - request: Any, - input_tokens: int = 0, - *, - request_id: str | None = None, - ) -> AsyncIterator[str]: - """Stream response via a native Anthropic-compatible messages endpoint.""" - tag = self._provider_name - req_tag = f" request_id={request_id}" if request_id else "" - thinking_enabled = self._is_thinking_enabled(request) - body = self._build_request_body(request) - - logger.info( - "{}_STREAM:{} natively passing Anthropic request model={} msgs={} tools={}", - tag, - req_tag, - body.get("model"), - len(body.get("messages", [])), - len(body.get("tools", [])), - ) - - response: httpx.Response | None = None - sent_any_event = False - state = self._new_stream_state(request, thinking_enabled=thinking_enabled) - - async with self._global_rate_limiter.concurrency_slot(): - try: - response = await self._global_rate_limiter.execute_with_retry( - self._send_stream_request, body - ) - - if response.status_code != 200: - await self._raise_for_status(response, req_tag=req_tag) - - async for chunk in self._iter_stream_chunks( - response, - state=state, - thinking_enabled=thinking_enabled, - ): - sent_any_event = True - yield chunk - - except Exception as error: - logger.error( - "{}_ERROR:{} {}: {}", tag, req_tag, type(error).__name__, error - ) - error_message = self._get_error_message(error, request_id) - - if response is not None and not response.is_closed: - await response.aclose() - - logger.info( - "{}_STREAM: Emitting native SSE error event for {}{}", - tag, - type(error).__name__, - req_tag, - ) - for event in self._emit_error_events( - request=request, - input_tokens=input_tokens, - error_message=error_message, - sent_any_event=sent_any_event, - ): - yield event - return - finally: - if response is not None and not response.is_closed: - await response.aclose() diff --git a/providers/llamacpp/client.py b/providers/llamacpp/client.py index 5ac9826f..5dfcf5bb 100644 --- a/providers/llamacpp/client.py +++ b/providers/llamacpp/client.py @@ -1,17 +1,159 @@ """Llama.cpp provider implementation.""" -from providers.anthropic_compat import AnthropicMessagesProvider -from providers.base import ProviderConfig +import json +from collections.abc import AsyncIterator +from typing import Any + +import httpx +from loguru import logger + +from providers.base import BaseProvider, ProviderConfig +from providers.common import get_user_facing_error_message, map_error +from providers.rate_limit import GlobalRateLimiter LLAMACPP_DEFAULT_BASE_URL = "http://localhost:8080/v1" -class LlamaCppProvider(AnthropicMessagesProvider): +class LlamaCppProvider(BaseProvider): """Llama.cpp provider using native Anthropic Messages API endpoint.""" def __init__(self, config: ProviderConfig): - super().__init__( - config, - provider_name="LLAMACPP", - default_base_url=LLAMACPP_DEFAULT_BASE_URL, + super().__init__(config) + self._provider_name = "LLAMACPP" + self._base_url = (config.base_url or LLAMACPP_DEFAULT_BASE_URL).rstrip("/") + + # We need the base URL without /v1 if the user provided it with /v1, + # so we can append /v1/messages safely. + # Actually, if they provided http://localhost:8080/v1, we can just use + # {base_url}/messages which becomes http://localhost:8080/v1/messages + + self._global_rate_limiter = GlobalRateLimiter.get_instance( + rate_limit=config.rate_limit, + rate_window=config.rate_window, + max_concurrency=config.max_concurrency, ) + self._client = httpx.AsyncClient( + base_url=self._base_url, + proxy=config.proxy or None, + timeout=httpx.Timeout( + config.http_read_timeout, + connect=config.http_connect_timeout, + read=config.http_read_timeout, + write=config.http_write_timeout, + ), + ) + + async def cleanup(self) -> None: + """Release HTTP client resources.""" + await self._client.aclose() + + async def stream_response( + self, + request: Any, + input_tokens: int = 0, + *, + request_id: str | None = None, + ) -> AsyncIterator[str]: + """Stream response natively via Llama.cpp's Anthropic-compatible endpoint.""" + tag = self._provider_name + req_tag = f" request_id={request_id}" if request_id else "" + thinking_enabled = self._is_thinking_enabled(request) + + # Dump the Anthropic Pydantic model directly into a dict + body = request.model_dump(exclude_none=True) + + # Remove extra_body, original_model, resolved_provider_model which are internal + body.pop("extra_body", None) + body.pop("original_model", None) + body.pop("resolved_provider_model", None) + + # Translate internal ThinkingConfig to Anthropic API schema + if "thinking" in body: + thinking_cfg = body.pop("thinking") + if ( + thinking_enabled + and isinstance(thinking_cfg, dict) + and thinking_cfg.get("enabled") + ): + # Anthropic API requires a budget_tokens value when enabled + body["thinking"] = {"type": "enabled"} + + # Ensure max_tokens is present (Claude API requires it) + if "max_tokens" not in body: + body["max_tokens"] = 81920 + + logger.info( + "{}_STREAM:{} natively passing Anthropic request to llama.cpp model={} msgs={} tools={}", + tag, + req_tag, + body.get("model"), + len(body.get("messages", [])), + len(body.get("tools", [])), + ) + + async with self._global_rate_limiter.concurrency_slot(): + try: + # We use execute_with_retry around the streaming request context + # To do this safely with httpx streaming, we await the chunk stream + + async def _make_request(): + request_obj = self._client.build_request( + "POST", + "/messages", + json=body, + headers={"Content-Type": "application/json"}, + ) + return await self._client.send(request_obj, stream=True) + + response = await self._global_rate_limiter.execute_with_retry( + _make_request + ) + + if response.status_code != 200: + try: + response.raise_for_status() + except httpx.HTTPStatusError as e: + text = await response.aread() + logger.error( + "{}_ERROR:{} HTTP {}: {}", + tag, + req_tag, + response.status_code, + text.decode("utf-8", errors="replace"), + ) + raise e + + async for line in response.aiter_lines(): + if line: + yield f"{line}\n" + else: + yield "\n" + + except Exception as e: + logger.error("{}_ERROR:{} {}: {}", tag, req_tag, type(e).__name__, e) + mapped_e = map_error(e) + if getattr(mapped_e, "status_code", None) == 405: + error_message = ( + f"Upstream provider {tag} rejected the request method " + "or endpoint (HTTP 405)." + ) + else: + error_message = get_user_facing_error_message( + mapped_e, read_timeout_s=self._config.http_read_timeout + ) + if request_id: + error_message += f"\nRequest ID: {request_id}" + + logger.info( + "{}_STREAM: Emitting native SSE error event for {}{}", + tag, + type(e).__name__, + req_tag, + ) + + # Emit an Anthropic-compatible error event + error_event = { + "type": "error", + "error": {"type": "api_error", "message": error_message}, + } + yield f"event: error\ndata: {json.dumps(error_event)}\n\n" diff --git a/providers/lmstudio/client.py b/providers/lmstudio/client.py index 29e227fa..c6d66f7f 100644 --- a/providers/lmstudio/client.py +++ b/providers/lmstudio/client.py @@ -1,17 +1,159 @@ """LM Studio provider implementation.""" -from providers.anthropic_compat import AnthropicMessagesProvider -from providers.base import ProviderConfig +import json +from collections.abc import AsyncIterator +from typing import Any + +import httpx +from loguru import logger + +from providers.base import BaseProvider, ProviderConfig +from providers.common import get_user_facing_error_message, map_error +from providers.rate_limit import GlobalRateLimiter LMSTUDIO_DEFAULT_BASE_URL = "http://localhost:1234/v1" -class LMStudioProvider(AnthropicMessagesProvider): +class LMStudioProvider(BaseProvider): """LM Studio provider using native Anthropic Messages API endpoint.""" def __init__(self, config: ProviderConfig): - super().__init__( - config, - provider_name="LMSTUDIO", - default_base_url=LMSTUDIO_DEFAULT_BASE_URL, + super().__init__(config) + self._provider_name = "LMSTUDIO" + self._base_url = (config.base_url or LMSTUDIO_DEFAULT_BASE_URL).rstrip("/") + + # We need the base URL without /v1 if the user provided it with /v1, + # so we can append /v1/messages safely. + # Actually, if they provided http://localhost:1234/v1, we can just use + # {base_url}/messages which becomes http://localhost:1234/v1/messages + + self._global_rate_limiter = GlobalRateLimiter.get_instance( + rate_limit=config.rate_limit, + rate_window=config.rate_window, + max_concurrency=config.max_concurrency, ) + self._client = httpx.AsyncClient( + base_url=self._base_url, + proxy=config.proxy or None, + timeout=httpx.Timeout( + config.http_read_timeout, + connect=config.http_connect_timeout, + read=config.http_read_timeout, + write=config.http_write_timeout, + ), + ) + + async def cleanup(self) -> None: + """Release HTTP client resources.""" + await self._client.aclose() + + async def stream_response( + self, + request: Any, + input_tokens: int = 0, + *, + request_id: str | None = None, + ) -> AsyncIterator[str]: + """Stream response natively via LM Studio's Anthropic-compatible endpoint.""" + tag = self._provider_name + req_tag = f" request_id={request_id}" if request_id else "" + thinking_enabled = self._is_thinking_enabled(request) + + # Dump the Anthropic Pydantic model directly into a dict + body = request.model_dump(exclude_none=True) + + # Remove extra_body, original_model, resolved_provider_model which are internal + body.pop("extra_body", None) + body.pop("original_model", None) + body.pop("resolved_provider_model", None) + + # Translate internal ThinkingConfig to Anthropic API schema + if "thinking" in body: + thinking_cfg = body.pop("thinking") + if ( + thinking_enabled + and isinstance(thinking_cfg, dict) + and thinking_cfg.get("enabled") + ): + # Anthropic API requires a budget_tokens value when enabled + body["thinking"] = {"type": "enabled"} + + # Ensure max_tokens is present (Claude API requires it) + if "max_tokens" not in body: + body["max_tokens"] = 81920 + + logger.info( + "{}_STREAM:{} natively passing Anthropic request to LMStudio model={} msgs={} tools={}", + tag, + req_tag, + body.get("model"), + len(body.get("messages", [])), + len(body.get("tools", [])), + ) + + async with self._global_rate_limiter.concurrency_slot(): + try: + # We use execute_with_retry around the streaming request context + # To do this safely with httpx streaming, we await the chunk stream + + async def _make_request(): + request_obj = self._client.build_request( + "POST", + "/messages", + json=body, + headers={"Content-Type": "application/json"}, + ) + return await self._client.send(request_obj, stream=True) + + response = await self._global_rate_limiter.execute_with_retry( + _make_request + ) + + if response.status_code != 200: + try: + response.raise_for_status() + except httpx.HTTPStatusError as e: + text = await response.aread() + logger.error( + "{}_ERROR:{} HTTP {}: {}", + tag, + req_tag, + response.status_code, + text.decode("utf-8", errors="replace"), + ) + raise e + + async for line in response.aiter_lines(): + if line: + yield f"{line}\n" + else: + yield "\n" + + except Exception as e: + logger.error("{}_ERROR:{} {}: {}", tag, req_tag, type(e).__name__, e) + mapped_e = map_error(e) + if getattr(mapped_e, "status_code", None) == 405: + error_message = ( + f"Upstream provider {tag} rejected the request method " + "or endpoint (HTTP 405)." + ) + else: + error_message = get_user_facing_error_message( + mapped_e, read_timeout_s=self._config.http_read_timeout + ) + if request_id: + error_message += f"\nRequest ID: {request_id}" + + logger.info( + "{}_STREAM: Emitting native SSE error event for {}{}", + tag, + type(e).__name__, + req_tag, + ) + + # Emit an Anthropic-compatible error event + error_event = { + "type": "error", + "error": {"type": "api_error", "message": error_message}, + } + yield f"event: error\ndata: {json.dumps(error_event)}\n\n" diff --git a/providers/open_router/client.py b/providers/open_router/client.py index a9603816..af04f222 100644 --- a/providers/open_router/client.py +++ b/providers/open_router/client.py @@ -1,176 +1,45 @@ """OpenRouter provider implementation.""" -import json -import uuid from collections.abc import Iterator -from dataclasses import dataclass, field from typing import Any -from providers.anthropic_compat import AnthropicMessagesProvider, StreamChunkMode from providers.base import ProviderConfig -from providers.common import SSEBuilder, append_request_id +from providers.common import SSEBuilder +from providers.openai_compat import OpenAICompatibleProvider from .request import build_request_body OPENROUTER_BASE_URL = "https://openrouter.ai/api/v1" -_ANTHROPIC_VERSION = "2023-06-01" -@dataclass -class _SSEFilterState: - """Track Anthropic content block index remapping while filtering thinking.""" - - next_index: int = 0 - index_map: dict[int, int] = field(default_factory=dict) - dropped_indexes: set[int] = field(default_factory=set) - - -class OpenRouterProvider(AnthropicMessagesProvider): - """OpenRouter provider using the Anthropic-compatible messages API.""" - - stream_chunk_mode: StreamChunkMode = "event" +class OpenRouterProvider(OpenAICompatibleProvider): + """OpenRouter provider using OpenAI-compatible API.""" def __init__(self, config: ProviderConfig): super().__init__( config, provider_name="OPENROUTER", - default_base_url=OPENROUTER_BASE_URL, + base_url=config.base_url or OPENROUTER_BASE_URL, + api_key=config.api_key, ) def _build_request_body(self, request: Any) -> dict: - """Internal helper for tests and direct request dispatch.""" + """Internal helper for tests and shared building.""" return build_request_body( request, thinking_enabled=self._is_thinking_enabled(request), ) - def _request_headers(self) -> dict[str, str]: - """Return OpenRouter's Anthropic-compatible messages headers.""" - return { - "Accept": "text/event-stream", - "Authorization": f"Bearer {self._api_key}", - "Content-Type": "application/json", - "anthropic-version": _ANTHROPIC_VERSION, - } - - @staticmethod - def _format_sse_event(event_name: str | None, data_text: str) -> str: - """Format an SSE event from its event name and data payload.""" - lines: list[str] = [] - if event_name: - lines.append(f"event: {event_name}") - lines.extend(f"data: {line}" for line in data_text.splitlines()) - return "\n".join(lines) + "\n\n" - - @staticmethod - def _parse_sse_event(event: str) -> tuple[str | None, str]: - """Extract the event name and raw data payload from an SSE event.""" - event_name = None - data_lines: list[str] = [] - for line in event.strip().splitlines(): - if line.startswith("event:"): - event_name = line[6:].strip() - elif line.startswith("data:"): - data_lines.append(line[5:].lstrip()) - return event_name, "\n".join(data_lines) - - @staticmethod - def _remap_index( - payload: dict[str, Any], state: _SSEFilterState, *, create: bool - ) -> int | None: - """Return the downstream index for a content block event.""" - upstream_index = payload.get("index") - if not isinstance(upstream_index, int): - return None - if upstream_index in state.dropped_indexes: - return None - mapped_index = state.index_map.get(upstream_index) - if mapped_index is None and create: - mapped_index = state.next_index - state.index_map[upstream_index] = mapped_index - state.next_index += 1 - return mapped_index - - def _filter_sse_event(self, event: str, state: _SSEFilterState) -> str | None: - """Drop upstream thinking blocks and remap the remaining block indexes.""" - event_name, data_text = self._parse_sse_event(event) - if not event_name or not data_text: - return event - - try: - payload = json.loads(data_text) - except json.JSONDecodeError: - return event - - if event_name == "content_block_start": - block = payload.get("content_block") - block_type = block.get("type") if isinstance(block, dict) else None - upstream_index = payload.get("index") - if isinstance(block_type, str) and "thinking" in block_type: - if isinstance(upstream_index, int): - state.dropped_indexes.add(upstream_index) - return None - - mapped_index = self._remap_index(payload, state, create=True) - if mapped_index is not None: - payload["index"] = mapped_index - return self._format_sse_event(event_name, json.dumps(payload)) - - if event_name == "content_block_delta": - delta = payload.get("delta") - delta_type = delta.get("type") if isinstance(delta, dict) else None - if isinstance(delta_type, str) and "thinking" in delta_type: - return None - - mapped_index = self._remap_index(payload, state, create=False) - if mapped_index is None: - return None - payload["index"] = mapped_index - return self._format_sse_event(event_name, json.dumps(payload)) - - if event_name == "content_block_stop": - mapped_index = self._remap_index(payload, state, create=False) - if mapped_index is None: - return None - payload["index"] = mapped_index - return self._format_sse_event(event_name, json.dumps(payload)) - - return event - - def _new_stream_state(self, request: Any, *, thinking_enabled: bool) -> Any: - """Create per-stream state for thinking block filtering.""" - return _SSEFilterState() - - def _transform_stream_event( - self, - event: str, - state: Any, - *, - thinking_enabled: bool, - ) -> str | None: - """Drop thinking events when thinking is disabled.""" - if thinking_enabled: - return event - if isinstance(state, _SSEFilterState): - return self._filter_sse_event(event, state) - return event - - def _format_error_message(self, base_message: str, request_id: str | None) -> str: - """Keep OpenRouter's existing request-id suffix format.""" - return append_request_id(base_message, request_id) - - def _emit_error_events( - self, - *, - request: Any, - input_tokens: int, - error_message: str, - sent_any_event: bool, + def _handle_extra_reasoning( + self, delta: Any, sse: SSEBuilder, *, thinking_enabled: bool ) -> Iterator[str]: - """Emit the existing Anthropic SSE error shape.""" - sse = SSEBuilder(f"msg_{uuid.uuid4()}", request.model, input_tokens) - if not sent_any_event: - yield sse.message_start() - yield from sse.emit_error(error_message) - yield sse.message_delta("end_turn", 1) - yield sse.message_stop() + """Handle reasoning_details for StepFun models.""" + if not thinking_enabled: + return + reasoning_details = getattr(delta, "reasoning_details", None) + if reasoning_details and isinstance(reasoning_details, list): + for item in reasoning_details: + text = item.get("text", "") if isinstance(item, dict) else "" + if text: + yield from sse.ensure_thinking_block() + yield sse.emit_thinking_delta(text) diff --git a/providers/open_router/request.py b/providers/open_router/request.py index 93a6c296..c524a006 100644 --- a/providers/open_router/request.py +++ b/providers/open_router/request.py @@ -1,101 +1,39 @@ """Request builder for OpenRouter provider.""" -from collections.abc import Sequence from typing import Any from loguru import logger -from pydantic import BaseModel + +from providers.common.message_converter import build_base_request_body OPENROUTER_DEFAULT_MAX_TOKENS = 81920 -_REQUEST_FIELDS = ( - "model", - "messages", - "system", - "max_tokens", - "stop_sequences", - "stream", - "temperature", - "top_p", - "top_k", - "metadata", - "tools", - "tool_choice", - "thinking", - "extra_body", - "original_model", - "resolved_provider_model", -) - -_INTERNAL_FIELDS = { - "thinking", - "extra_body", - "original_model", - "resolved_provider_model", -} - - -def _serialize_value(value: Any) -> Any: - """Convert Pydantic models and lightweight objects into JSON-ready values.""" - if isinstance(value, BaseModel): - return value.model_dump(exclude_none=True) - if isinstance(value, dict): - return { - key: _serialize_value(item) - for key, item in value.items() - if item is not None - } - if isinstance(value, Sequence) and not isinstance(value, str | bytes | bytearray): - return [_serialize_value(item) for item in value] - if value is None or isinstance(value, str | int | float | bool): - return value - if hasattr(value, "__dict__"): - return { - key: _serialize_value(item) - for key, item in vars(value).items() - if not key.startswith("_") and item is not None - } - return value - - -def _dump_request_fields(request_data: Any) -> dict[str, Any]: - """Extract the public request fields we forward to OpenRouter.""" - if isinstance(request_data, BaseModel): - return request_data.model_dump(exclude_none=True) - - dumped: dict[str, Any] = {} - for field in _REQUEST_FIELDS: - value = getattr(request_data, field, None) - if value is not None: - dumped[field] = _serialize_value(value) - return dumped - def build_request_body(request_data: Any, *, thinking_enabled: bool) -> dict: - """Build an Anthropic-format request body for OpenRouter's messages API.""" + """Build OpenAI-format request body from Anthropic request for OpenRouter.""" logger.debug( "OPENROUTER_REQUEST: conversion start model={} msgs={}", getattr(request_data, "model", "?"), len(getattr(request_data, "messages", [])), ) + body = build_base_request_body( + request_data, + include_thinking=thinking_enabled, + default_max_tokens=OPENROUTER_DEFAULT_MAX_TOKENS, + include_reasoning_for_openrouter=thinking_enabled, + ) - dumped_request = _dump_request_fields(request_data) - request_extra = dumped_request.pop("extra_body", None) - body = { - key: value - for key, value in dumped_request.items() - if key not in _INTERNAL_FIELDS - } - - if isinstance(request_extra, dict): - body.update(request_extra) - - body["stream"] = True - if body.get("max_tokens") is None: - body["max_tokens"] = OPENROUTER_DEFAULT_MAX_TOKENS + # OpenRouter reasoning: extra_body={"reasoning": {"enabled": True}} + extra_body: dict[str, Any] = {} + request_extra = getattr(request_data, "extra_body", None) + if request_extra: + extra_body.update(request_extra) if thinking_enabled: - body.setdefault("reasoning", {"enabled": True}) + extra_body.setdefault("reasoning", {"enabled": True}) + + if extra_body: + body["extra_body"] = extra_body logger.debug( "OPENROUTER_REQUEST: conversion done model={} msgs={} tools={}", diff --git a/tests/providers/test_anthropic_compat.py b/tests/providers/test_anthropic_compat.py deleted file mode 100644 index 4f84c381..00000000 --- a/tests/providers/test_anthropic_compat.py +++ /dev/null @@ -1,203 +0,0 @@ -"""Tests for the shared native Anthropic Messages provider base.""" - -from contextlib import asynccontextmanager -from unittest.mock import AsyncMock, MagicMock, patch - -import httpx -import pytest - -from providers.anthropic_compat import AnthropicMessagesProvider -from providers.base import ProviderConfig - - -class NativeProvider(AnthropicMessagesProvider): - """Concrete provider used to exercise the shared base.""" - - def __init__(self, config: ProviderConfig): - super().__init__( - config, - provider_name="TEST_NATIVE", - default_base_url="https://example.test/v1", - ) - - def _request_headers(self) -> dict[str, str]: - return {"Content-Type": "application/json", "X-Test": "1"} - - -class MockRequest: - model = "test-model" - - def __init__(self, *, thinking_enabled: bool = True, body: dict | None = None): - self.thinking = MagicMock() - self.thinking.enabled = thinking_enabled - self._body = body or { - "model": self.model, - "messages": [{"role": "user", "content": "Hello"}], - "extra_body": {"ignored": True}, - "original_model": "claude", - "resolved_provider_model": "native/test-model", - "thinking": {"enabled": thinking_enabled}, - } - - def model_dump(self, exclude_none=True): - return dict(self._body) - - -class FakeResponse: - def __init__(self, *, status_code=200, lines=None, text=""): - self.status_code = status_code - self._lines = lines or [] - self._text = text - self.is_closed = False - self.request = httpx.Request("POST", "https://example.test/v1/messages") - - async def aiter_lines(self): - for line in self._lines: - yield line - - async def aread(self): - return self._text.encode() - - def raise_for_status(self): - response = httpx.Response( - self.status_code, - request=self.request, - text=self._text, - ) - response.raise_for_status() - - async def aclose(self): - self.is_closed = True - - -@pytest.fixture -def provider_config(): - return ProviderConfig( - api_key="test-key", - base_url="https://custom.test/v1/", - proxy="socks5://127.0.0.1:9999", - rate_limit=10, - rate_window=60, - http_read_timeout=600.0, - http_write_timeout=15.0, - http_connect_timeout=5.0, - ) - - -@pytest.fixture(autouse=True) -def mock_rate_limiter(): - @asynccontextmanager - async def _slot(): - yield - - with patch("providers.anthropic_compat.GlobalRateLimiter") as mock: - instance = mock.get_instance.return_value - - async def _passthrough(fn, *args, **kwargs): - return await fn(*args, **kwargs) - - instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot - yield instance - - -def test_init_configures_httpx_client(provider_config): - with patch("httpx.AsyncClient") as mock_client: - provider = NativeProvider(provider_config) - - assert provider._provider_name == "TEST_NATIVE" - assert provider._api_key == "test-key" - assert provider._base_url == "https://custom.test/v1" - kwargs = mock_client.call_args.kwargs - timeout = kwargs["timeout"] - assert kwargs["base_url"] == "https://custom.test/v1" - assert kwargs["proxy"] == "socks5://127.0.0.1:9999" - assert timeout.read == 600.0 - assert timeout.write == 15.0 - assert timeout.connect == 5.0 - - -def test_default_request_body_strips_internal_fields(provider_config): - provider = NativeProvider(provider_config) - - body = provider._build_request_body(MockRequest()) - - assert body["model"] == "test-model" - assert body["thinking"] == {"type": "enabled"} - assert body["max_tokens"] == 81920 - assert "extra_body" not in body - assert "original_model" not in body - assert "resolved_provider_model" not in body - - -@pytest.mark.asyncio -async def test_stream_uses_retry_builds_request_and_closes_response( - provider_config, - mock_rate_limiter, -): - provider = NativeProvider(provider_config) - req = MockRequest() - request_obj = httpx.Request("POST", "https://custom.test/v1/messages") - response = FakeResponse( - lines=[ - "event: message_start", - 'data: {"type":"message_start"}', - "", - ] - ) - - with ( - patch.object( - provider._client, "build_request", return_value=request_obj - ) as mock_build, - patch.object( - provider._client, - "send", - new_callable=AsyncMock, - return_value=response, - ) as mock_send, - ): - events = [event async for event in provider.stream_response(req)] - - assert events == [ - "event: message_start\n", - 'data: {"type":"message_start"}\n', - "\n", - ] - assert response.is_closed - assert mock_build.call_args.args[:2] == ("POST", "/messages") - assert mock_build.call_args.kwargs["headers"] == { - "Content-Type": "application/json", - "X-Test": "1", - } - assert mock_build.call_args.kwargs["json"]["thinking"] == {"type": "enabled"} - mock_send.assert_awaited_once_with(request_obj, stream=True) - mock_rate_limiter.execute_with_retry.assert_awaited_once() - - -@pytest.mark.asyncio -async def test_stream_maps_non_200_to_error_event_and_closes_response( - provider_config, -): - provider = NativeProvider(provider_config) - req = MockRequest() - response = FakeResponse(status_code=500, text="Internal Server Error") - - with ( - patch.object(provider._client, "build_request", return_value=MagicMock()), - patch.object( - provider._client, - "send", - new_callable=AsyncMock, - return_value=response, - ), - ): - events = [ - event async for event in provider.stream_response(req, request_id="REQ_123") - ] - - assert response.is_closed - assert len(events) == 1 - assert events[0].startswith("event: error\ndata: {") - assert "Internal Server Error" in events[0] - assert "REQ_123" in events[0] diff --git a/tests/providers/test_llamacpp.py b/tests/providers/test_llamacpp.py index 28a34c5e..e1a9f05b 100644 --- a/tests/providers/test_llamacpp.py +++ b/tests/providers/test_llamacpp.py @@ -1,6 +1,5 @@ """Tests for Llama.cpp native Anthropic provider.""" -from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch import httpx @@ -56,12 +55,7 @@ def llamacpp_config(): @pytest.fixture(autouse=True) def mock_rate_limiter(): """Mock the global rate limiter to prevent waiting.""" - - @asynccontextmanager - async def _slot(): - yield - - with patch("providers.anthropic_compat.GlobalRateLimiter") as mock: + with patch("providers.llamacpp.client.GlobalRateLimiter") as mock: instance = mock.get_instance.return_value instance.wait_if_blocked = AsyncMock(return_value=False) @@ -69,7 +63,6 @@ def mock_rate_limiter(): return await fn(*args, **kwargs) instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot yield instance diff --git a/tests/providers/test_lmstudio.py b/tests/providers/test_lmstudio.py index a0dbe674..1173846c 100644 --- a/tests/providers/test_lmstudio.py +++ b/tests/providers/test_lmstudio.py @@ -1,6 +1,5 @@ """Tests for LM Studio native Anthropic provider.""" -from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch import httpx @@ -56,12 +55,7 @@ def lmstudio_config(): @pytest.fixture(autouse=True) def mock_rate_limiter(): """Mock the global rate limiter to prevent waiting.""" - - @asynccontextmanager - async def _slot(): - yield - - with patch("providers.anthropic_compat.GlobalRateLimiter") as mock: + with patch("providers.lmstudio.client.GlobalRateLimiter") as mock: instance = mock.get_instance.return_value instance.wait_if_blocked = AsyncMock(return_value=False) @@ -69,7 +63,6 @@ def mock_rate_limiter(): return await fn(*args, **kwargs) instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.concurrency_slot.side_effect = _slot yield instance diff --git a/tests/providers/test_open_router.py b/tests/providers/test_open_router.py index cec13679..fbcfd475 100644 --- a/tests/providers/test_open_router.py +++ b/tests/providers/test_open_router.py @@ -1,10 +1,8 @@ -"""Tests for the OpenRouter provider.""" +"""Tests for OpenRouter provider.""" import json -from contextlib import asynccontextmanager from unittest.mock import AsyncMock, MagicMock, patch -import httpx import pytest from providers.base import ProviderConfig @@ -18,16 +16,6 @@ class MockMessage: self.content = content -class MockTool: - def __init__(self): - self.name = "run_command" - self.description = "Run a command" - self.input_schema = { - "type": "object", - "properties": {"cmd": {"type": "string"}}, - } - - class MockRequest: def __init__(self, **kwargs): self.model = "stepfun/step-3.5-flash:free" @@ -35,58 +23,14 @@ class MockRequest: self.max_tokens = 100 self.temperature = 0.5 self.top_p = 0.9 - self.top_k = 20 self.system = "System prompt" - self.stop_sequences = ["STOP"] - self.stream = False - self.metadata = {"source": "request"} - self.tools = [MockTool()] - self.tool_choice = {"type": "auto"} + self.stop_sequences = None + self.tools = [] self.extra_body = {} - self.original_model = "claude-sonnet-4-20250514" - self.resolved_provider_model = "open_router/stepfun/step-3.5-flash:free" self.thinking = MagicMock() self.thinking.enabled = True - for key, value in kwargs.items(): - setattr(self, key, value) - - -class FakeResponse: - """Simple async streaming response for provider tests.""" - - def __init__(self, *, status_code=200, lines=None, text=""): - self.status_code = status_code - self._lines = lines or [] - self._text = text - self.is_closed = False - self.request = httpx.Request("POST", "https://openrouter.ai/api/v1/messages") - - async def aiter_lines(self): - for line in self._lines: - yield line - - def raise_for_status(self): - response = httpx.Response( - self.status_code, - request=self.request, - text=self._text, - ) - response.raise_for_status() - - async def aclose(self): - self.is_closed = True - - -def parse_sse_event(event: str) -> tuple[str | None, dict]: - """Parse an SSE event string into event type and JSON payload.""" - event_type = None - data_lines: list[str] = [] - for line in event.strip().splitlines(): - if line.startswith("event:"): - event_type = line[6:].strip() - elif line.startswith("data:"): - data_lines.append(line[5:].lstrip()) - return event_type, json.loads("\n".join(data_lines)) + for k, v in kwargs.items(): + setattr(self, k, v) @pytest.fixture @@ -102,20 +46,14 @@ def open_router_config(): @pytest.fixture(autouse=True) def mock_rate_limiter(): """Mock the global rate limiter to prevent waiting.""" - - @asynccontextmanager - async def _slot(): - yield - - with patch("providers.anthropic_compat.GlobalRateLimiter") as mock: + with patch("providers.openai_compat.GlobalRateLimiter") as mock: instance = mock.get_instance.return_value + instance.wait_if_blocked = AsyncMock(return_value=False) async def _passthrough(fn, *args, **kwargs): return await fn(*args, **kwargs) instance.execute_with_retry = AsyncMock(side_effect=_passthrough) - instance.wait_if_blocked = AsyncMock(return_value=False) - instance.concurrency_slot.side_effect = _slot yield instance @@ -124,75 +62,47 @@ def open_router_provider(open_router_config): return OpenRouterProvider(open_router_config) -def test_init_uses_httpx_client_with_proxy_and_timeouts(): - """Provider initialization configures an httpx client directly.""" +def test_init(open_router_config): + """Test provider initialization.""" + with patch("providers.openai_compat.AsyncOpenAI") as mock_openai: + provider = OpenRouterProvider(open_router_config) + assert provider._api_key == "test_openrouter_key" + assert provider._base_url == "https://openrouter.ai/api/v1" + mock_openai.assert_called_once() + + +def test_init_uses_configurable_timeouts(): + """Test that provider passes configurable read/write/connect timeouts to client.""" config = ProviderConfig( api_key="test_openrouter_key", base_url="https://openrouter.ai/api/v1", - proxy="socks5://127.0.0.1:9999", http_read_timeout=600.0, http_write_timeout=15.0, http_connect_timeout=5.0, ) - - with patch("httpx.AsyncClient") as mock_client: - provider = OpenRouterProvider(config) - - assert provider._api_key == "test_openrouter_key" - assert provider._base_url == "https://openrouter.ai/api/v1" - kwargs = mock_client.call_args.kwargs - timeout = kwargs["timeout"] - assert kwargs["base_url"] == "https://openrouter.ai/api/v1" - assert kwargs["proxy"] == "socks5://127.0.0.1:9999" - assert timeout.read == 600.0 - assert timeout.write == 15.0 - assert timeout.connect == 5.0 + with patch("providers.openai_compat.AsyncOpenAI") as mock_openai: + OpenRouterProvider(config) + call_kwargs = mock_openai.call_args[1] + timeout = call_kwargs["timeout"] + assert timeout.read == 600.0 + assert timeout.write == 15.0 + assert timeout.connect == 5.0 -def test_build_request_body_is_anthropic_shaped(open_router_provider): - """System stays top-level and internal fields are stripped.""" +def test_build_request_body_has_reasoning_extra(open_router_provider): + """Request body has extra_body.reasoning.enabled for thinking models.""" req = MockRequest() - body = open_router_provider._build_request_body(req) assert body["model"] == "stepfun/step-3.5-flash:free" - assert body["system"] == "System prompt" - assert body["stream"] is True assert body["temperature"] == 0.5 - assert body["top_p"] == 0.9 - assert body["top_k"] == 20 - assert body["stop_sequences"] == ["STOP"] - assert body["metadata"] == {"source": "request"} - assert body["tool_choice"] == {"type": "auto"} - assert len(body["messages"]) == 1 - assert body["messages"][0] == {"role": "user", "content": "Hello"} - assert body["tools"][0]["name"] == "run_command" - assert body["reasoning"] == {"enabled": True} - assert "thinking" not in body - assert "original_model" not in body - assert "resolved_provider_model" not in body - assert "extra_body" not in body + assert len(body["messages"]) == 2 # System + User + assert body["messages"][0]["role"] == "system" + assert body["messages"][0]["content"] == "System prompt" - -def test_build_request_body_extra_body_merges_top_level_and_preserves_overrides( - open_router_provider, -): - req = MockRequest( - metadata={"source": "request"}, - extra_body={ - "metadata": {"source": "extra"}, - "reasoning": {"enabled": False}, - "service_tier": "flex", - "stream": False, - }, - ) - - body = open_router_provider._build_request_body(req) - - assert body["metadata"] == {"source": "extra"} - assert body["reasoning"] == {"enabled": False} - assert body["service_tier"] == "flex" - assert body["stream"] is True + assert "extra_body" in body + assert "reasoning" in body["extra_body"] + assert body["extra_body"]["reasoning"]["enabled"] is True def test_build_request_body_omits_reasoning_when_globally_disabled(open_router_config): @@ -200,10 +110,9 @@ def test_build_request_body_omits_reasoning_when_globally_disabled(open_router_c open_router_config.model_copy(update={"enable_thinking": False}) ) req = MockRequest() - body = provider._build_request_body(req) - assert "reasoning" not in body + assert "extra_body" not in body or "reasoning" not in body["extra_body"] def test_build_request_body_omits_reasoning_when_request_disables_thinking( @@ -211,174 +120,283 @@ def test_build_request_body_omits_reasoning_when_request_disables_thinking( ): req = MockRequest() req.thinking.enabled = False - body = open_router_provider._build_request_body(req) - assert "reasoning" not in body + assert "extra_body" not in body or "reasoning" not in body["extra_body"] + + +def test_build_request_body_base_url_and_model(open_router_provider): + """Base URL and model are correct in provider config.""" + assert open_router_provider._base_url == "https://openrouter.ai/api/v1" + req = MockRequest(model="stepfun/step-3.5-flash:free") + body = open_router_provider._build_request_body(req) + assert body["model"] == "stepfun/step-3.5-flash:free" def test_build_request_body_default_max_tokens(open_router_provider): + """max_tokens=None uses OPENROUTER_DEFAULT_MAX_TOKENS (81920).""" req = MockRequest(max_tokens=None) - body = open_router_provider._build_request_body(req) - assert body["max_tokens"] == OPENROUTER_DEFAULT_MAX_TOKENS + assert body["max_tokens"] == 81920 @pytest.mark.asyncio -async def test_stream_response_passthroughs_anthropic_sse(open_router_provider): +async def test_stream_response_text(open_router_provider): + """Test streaming text response.""" req = MockRequest() - request_obj = httpx.Request("POST", "https://openrouter.ai/api/v1/messages") - response = FakeResponse( - lines=[ - "event: message_start", - 'data: {"type":"message_start","message":{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"stepfun/step-3.5-flash:free","stop_reason":null,"stop_sequence":null,"usage":{"input_tokens":0,"output_tokens":0}}}', - "", - "event: content_block_start", - 'data: {"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}', - "", - "event: content_block_delta", - 'data: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello"}}', - "", - "event: content_block_stop", - 'data: {"type":"content_block_stop","index":0}', - "", - "event: message_delta", - 'data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{"input_tokens":22,"output_tokens":1}}', - "", - "event: message_stop", - 'data: {"type":"message_stop"}', - "", - ] - ) - with ( - patch.object( - open_router_provider._client, - "build_request", - return_value=request_obj, - ) as mock_build_request, - patch.object( - open_router_provider._client, - "send", - new_callable=AsyncMock, - return_value=response, - ) as mock_send, - ): + mock_chunk1 = MagicMock() + mock_chunk1.choices = [ + MagicMock( + delta=MagicMock(content="Hello", reasoning_content=None), + finish_reason=None, + ) + ] + mock_chunk1.usage = None + + mock_chunk2 = MagicMock() + mock_chunk2.choices = [ + MagicMock( + delta=MagicMock(content=" World", reasoning_content=None), + finish_reason="stop", + ) + ] + mock_chunk2.usage = MagicMock(completion_tokens=10) + + async def mock_stream(): + yield mock_chunk1 + yield mock_chunk2 + + with patch.object( + open_router_provider._client.chat.completions, "create", new_callable=AsyncMock + ) as mock_create: + mock_create.return_value = mock_stream() + events = [e async for e in open_router_provider.stream_response(req)] - assert events[0].startswith("event: message_start") - assert any("Hello" in event for event in events) - assert any("event: message_stop" in event for event in events) - assert mock_build_request.call_args.args[:2] == ("POST", "/messages") - assert ( - mock_build_request.call_args.kwargs["headers"]["anthropic-version"] - == "2023-06-01" - ) - assert mock_build_request.call_args.kwargs["headers"]["Authorization"].startswith( - "Bearer " - ) - mock_send.assert_awaited_once_with(request_obj, stream=True) + assert len(events) > 0 + assert "event: message_start" in events[0] + + text_content = "" + for e in events: + if "event: content_block_delta" in e and '"text_delta"' in e: + for line in e.splitlines(): + if line.startswith("data: "): + data = json.loads(line[6:]) + if "delta" in data and "text" in data["delta"]: + text_content += data["delta"]["text"] + + assert "Hello World" in text_content @pytest.mark.asyncio -async def test_stream_response_filters_thinking_when_disabled(open_router_config): +async def test_stream_response_reasoning_content(open_router_provider): + """Test streaming with reasoning_content delta.""" + req = MockRequest() + + mock_chunk = MagicMock() + mock_chunk.choices = [ + MagicMock( + delta=MagicMock(content=None, reasoning_content="Thinking..."), + finish_reason=None, + ) + ] + mock_chunk.usage = None + + async def mock_stream(): + yield mock_chunk + + with patch.object( + open_router_provider._client.chat.completions, "create", new_callable=AsyncMock + ) as mock_create: + mock_create.return_value = mock_stream() + + events = [e async for e in open_router_provider.stream_response(req)] + + found_thinking = False + for e in events: + if ( + "event: content_block_delta" in e + and '"thinking_delta"' in e + and "Thinking..." in e + ): + found_thinking = True + assert found_thinking + + +@pytest.mark.asyncio +async def test_stream_response_suppresses_reasoning_when_disabled(open_router_config): provider = OpenRouterProvider( open_router_config.model_copy(update={"enable_thinking": False}) ) req = MockRequest() - request_obj = httpx.Request("POST", "https://openrouter.ai/api/v1/messages") - response = FakeResponse( - lines=[ - "event: message_start", - 'data: {"type":"message_start","message":{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"stepfun/step-3.5-flash:free","stop_reason":null,"stop_sequence":null,"usage":{"input_tokens":0,"output_tokens":0}}}', - "", - "event: content_block_start", - 'data: {"type":"content_block_start","index":0,"content_block":{"type":"thinking","thinking":""}}', - "", - "event: content_block_delta", - 'data: {"type":"content_block_delta","index":0,"delta":{"type":"thinking_delta","thinking":"secret"}}', - "", - "event: content_block_stop", - 'data: {"type":"content_block_stop","index":0}', - "", - "event: content_block_start", - 'data: {"type":"content_block_start","index":1,"content_block":{"type":"text","text":""}}', - "", - "event: content_block_delta", - 'data: {"type":"content_block_delta","index":1,"delta":{"type":"text_delta","text":"Visible"}}', - "", - "event: content_block_stop", - 'data: {"type":"content_block_stop","index":1}', - "", - "event: message_delta", - 'data: {"type":"message_delta","delta":{"stop_reason":"end_turn","stop_sequence":null},"usage":{"input_tokens":22,"output_tokens":1}}', - "", - "event: message_stop", - 'data: {"type":"message_stop"}', - "", - ] - ) - with ( - patch.object(provider._client, "build_request", return_value=request_obj), - patch.object( - provider._client, - "send", - new_callable=AsyncMock, - return_value=response, - ), - ): + mock_chunk = MagicMock() + mock_chunk.choices = [ + MagicMock( + delta=MagicMock( + content="secretAnswer", + reasoning_content="Thinking...", + reasoning_details=[{"text": "Step 1"}], + ), + finish_reason="stop", + ) + ] + mock_chunk.usage = None + + async def mock_stream(): + yield mock_chunk + + with patch.object( + provider._client.chat.completions, "create", new_callable=AsyncMock + ) as mock_create: + mock_create.return_value = mock_stream() + events = [e async for e in provider.stream_response(req)] event_text = "".join(events) assert "thinking_delta" not in event_text + assert "Thinking..." not in event_text + assert "Step 1" not in event_text assert "secret" not in event_text - assert "Visible" in event_text - - content_block_starts = [ - parse_sse_event(event) - for event in events - if event.startswith("event: content_block_start") - ] - assert len(content_block_starts) == 1 - _, payload = content_block_starts[0] - assert payload["content_block"]["type"] == "text" - assert payload["index"] == 0 + assert "Answer" in event_text @pytest.mark.asyncio -async def test_stream_response_error_path_emits_existing_sse_shape( - open_router_provider, -): +async def test_stream_response_empty_choices_skipped(open_router_provider): + """Chunks with empty choices are skipped.""" req = MockRequest() - request_obj = httpx.Request("POST", "https://openrouter.ai/api/v1/messages") - response = FakeResponse(status_code=400, text='{"error":"bad request"}') - with ( - patch.object( - open_router_provider._client, - "build_request", - return_value=request_obj, - ), - patch.object( - open_router_provider._client, - "send", - new_callable=AsyncMock, - return_value=response, - ), - ): - events = [ - event - async for event in open_router_provider.stream_response( - req, - input_tokens=12, - request_id="req_123", - ) - ] + async def mock_stream(): + yield MagicMock(choices=[], usage=None) + yield MagicMock( + choices=[ + MagicMock( + delta=MagicMock(content="ok", reasoning_content=None), + finish_reason="stop", + ) + ], + usage=MagicMock(completion_tokens=2), + ) - assert events[0].startswith("event: message_start") - assert any("400 Bad Request" in event for event in events) - assert any("(request_id=req_123)" in event for event in events) - assert any("event: message_delta" in event for event in events) - assert any("event: message_stop" in event for event in events) + with patch.object( + open_router_provider._client.chat.completions, "create", new_callable=AsyncMock + ) as mock_create: + mock_create.return_value = mock_stream() + events = [e async for e in open_router_provider.stream_response(req)] + assert any("content_block_delta" in e and "ok" in e for e in events) + + +@pytest.mark.asyncio +async def test_stream_response_delta_none_skipped(open_router_provider): + """Chunks with delta=None are skipped.""" + req = MockRequest() + + async def mock_stream(): + yield MagicMock( + choices=[MagicMock(delta=None, finish_reason=None)], + usage=None, + ) + yield MagicMock( + choices=[ + MagicMock( + delta=MagicMock(content="x", reasoning_content=None), + finish_reason="stop", + ) + ], + usage=MagicMock(completion_tokens=1), + ) + + with patch.object( + open_router_provider._client.chat.completions, "create", new_callable=AsyncMock + ) as mock_create: + mock_create.return_value = mock_stream() + events = [e async for e in open_router_provider.stream_response(req)] + assert any("x" in e for e in events) + + +@pytest.mark.asyncio +async def test_stream_response_reasoning_details(open_router_provider): + """Streaming with reasoning_details (stepfun format).""" + req = MockRequest() + + mock_chunk = MagicMock() + mock_chunk.choices = [ + MagicMock( + delta=MagicMock( + content=None, + reasoning_content=None, + reasoning_details=[{"text": "Step 1"}], + ), + finish_reason=None, + ) + ] + mock_chunk.usage = None + + async def mock_stream(): + yield mock_chunk + yield MagicMock( + choices=[ + MagicMock( + delta=MagicMock( + content=None, + reasoning_content=None, + reasoning_details=None, + ), + finish_reason="stop", + ) + ], + usage=MagicMock(completion_tokens=5), + ) + + with patch.object( + open_router_provider._client.chat.completions, "create", new_callable=AsyncMock + ) as mock_create: + mock_create.return_value = mock_stream() + events = [e async for e in open_router_provider.stream_response(req)] + assert any("Step 1" in e for e in events) + + +@pytest.mark.asyncio +async def test_stream_response_error_path(open_router_provider): + """Stream raises exception -> error event emitted.""" + req = MockRequest() + + async def mock_stream(): + raise RuntimeError("API failed") + yield # unreachable, makes it a generator + + with patch.object( + open_router_provider._client.chat.completions, "create", new_callable=AsyncMock + ) as mock_create: + mock_create.return_value = mock_stream() + events = [e async for e in open_router_provider.stream_response(req)] + # Error is emitted; message_stop/done indicates stream completed + assert any("API failed" in e for e in events) + assert any("message_stop" in e for e in events) + + +@pytest.mark.asyncio +async def test_stream_response_finish_reason_only(open_router_provider): + """Chunk with finish_reason but no content still completes.""" + req = MockRequest() + + async def mock_stream(): + yield MagicMock( + choices=[ + MagicMock( + delta=MagicMock(content=None, reasoning_content=None), + finish_reason="stop", + ) + ], + usage=MagicMock(completion_tokens=0), + ) + + with patch.object( + open_router_provider._client.chat.completions, "create", new_callable=AsyncMock + ) as mock_create: + mock_create.return_value = mock_stream() + events = [e async for e in open_router_provider.stream_response(req)] + assert any("message_delta" in e for e in events) + assert any("message_stop" in e for e in events)