mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(ollama): one stream id, 400 on non-text tool content, separated tool results, no empty message item (#45177)
* fix(ollama): one stream id, 400 on non-text tool content, separated tool results, no empty message item The ollama/ text stream minted a fresh chatcmpl- id per chunk; every chunk now carries the one id of its response. A tool message with int content raised a TypeError inside ollama_pt and answered 500; it now answers 400 naming the message index and the type it carried. Merged user, tool, and function turns and the text parts of one message were joined with no separator; they are joined with a newline. A streamed tool call with no text rebuilds to content null instead of an empty string, and the Responses bridge no longer opens a message item on a leading empty delta, so a /v1/responses stream on a tool call completes with the function_call alone, as the non-stream call does * fix(ollama): defer the BadRequestError annotation so litellm imports on Python 3.10 to 3.13 * fix(ollama): read the user run lazily, type the prompt helpers, and 400 on a non-string text part * fix(ollama): 400 on a text or image_url part that carries no value * fix(responses): stream bridged output items at contiguous indexes and reject non-object ollama content parts The chat-completions bridge reserved output_index 0 for the message item, so a tool-only stream announced its function_call at index 1 with no item at 0 and the OpenAI SDK's responses.stream() accumulator raised IndexError. Items now take the next index in emission order, reasoning included, and response.completed lists the output in that streamed order. ollama_pt answers 400 for a content part that is not an object and for an image_url object without a url string, instead of a 500 or a silent drop. * fix(responses): list only announced items in the bridged completed snapshot The Responses bridge drops a message item from response.completed when the stream never announced one and its text is empty, so a tool-only stream and a reasoning-only stream end with the items the client saw. The chunk builder's rebuilt content goes back to what main does; aligning it with the non-stream shape is its own change for every provider's consumers * style(responses): sort the itertools import with the other stdlib imports * test(integration): add ollama prompt-tools and responses bridge audit cells
This commit is contained in:
parent
58258409c9
commit
ed9c02d2d9
9 changed files with 1293 additions and 61 deletions
|
|
@ -206,36 +206,108 @@ def _handle_ollama_system_message(messages: list, prompt: str, msg_i: int) -> tu
|
|||
return system_content_str, msg_i
|
||||
|
||||
|
||||
_OLLAMA_USER_ROLES: Final = frozenset({"user", "tool", "function"})
|
||||
|
||||
|
||||
def _ollama_bad_message(model: str, message: AllMessageValues, msg_i: int, detail: str) -> "litellm.BadRequestError":
|
||||
return litellm.BadRequestError(
|
||||
message=BAD_MESSAGE_ERROR_STR + f"the {message['role']} message at index {msg_i} {detail}",
|
||||
model=model,
|
||||
llm_provider="ollama",
|
||||
)
|
||||
|
||||
|
||||
def _ollama_content_part(model: str, message: AllMessageValues, part: object, msg_i: int) -> tuple[str, str]:
|
||||
match part:
|
||||
case {"type": "text", "text": str() as text}:
|
||||
return text, ""
|
||||
case {"type": "text", "text": bad_text}:
|
||||
raise _ollama_bad_message(
|
||||
model, message, msg_i, f"has a {type(bad_text).__name__} text part; text must be a string"
|
||||
)
|
||||
case {"type": "text"}:
|
||||
raise _ollama_bad_message(model, message, msg_i, "has a text part with no text; text must be a string")
|
||||
case {"type": "image_url", "image_url": str() as image_url}:
|
||||
return "", image_url
|
||||
case {"type": "image_url", "image_url": {"url": str() as image_url}}:
|
||||
return "", image_url
|
||||
case {"type": "image_url", "image_url": dict()}:
|
||||
raise _ollama_bad_message(
|
||||
model,
|
||||
message,
|
||||
msg_i,
|
||||
"has an image_url object without a url string; image_url must be a URL string or an object with a url",
|
||||
)
|
||||
case {"type": "image_url", "image_url": bad_image_url}:
|
||||
raise _ollama_bad_message(
|
||||
model,
|
||||
message,
|
||||
msg_i,
|
||||
f"has a {type(bad_image_url).__name__} image_url; image_url must be a URL string or an object with a url",
|
||||
)
|
||||
case {"type": "image_url"}:
|
||||
raise _ollama_bad_message(
|
||||
model,
|
||||
message,
|
||||
msg_i,
|
||||
"has an image_url part with no image_url; image_url must be a URL string or an object with a url",
|
||||
)
|
||||
case Mapping():
|
||||
return "", ""
|
||||
case _:
|
||||
raise _ollama_bad_message(
|
||||
model, message, msg_i, f"has a {type(part).__name__} content part; content parts must be objects"
|
||||
)
|
||||
|
||||
|
||||
def _ollama_user_message_parts(
|
||||
model: str, message: AllMessageValues, msg_i: int
|
||||
) -> tuple[tuple[str, ...], tuple[str, ...]]:
|
||||
msg_content: Final = message.get("content")
|
||||
if msg_content is None:
|
||||
return (), ()
|
||||
if isinstance(msg_content, str):
|
||||
return ((msg_content,) if msg_content else ()), ()
|
||||
if not isinstance(msg_content, list):
|
||||
raise _ollama_bad_message(
|
||||
model,
|
||||
message,
|
||||
msg_i,
|
||||
f"has {type(msg_content).__name__} content; content must be a string or a list of content parts",
|
||||
)
|
||||
parts: Final = tuple(_ollama_content_part(model, message, part, msg_i) for part in msg_content)
|
||||
texts: Final = tuple(text for text, _ in parts if text)
|
||||
image_urls: Final = tuple(image_url for _, image_url in parts if image_url)
|
||||
return texts, image_urls
|
||||
|
||||
|
||||
def _ollama_user_turn(model: str, messages: Sequence[AllMessageValues], msg_i: int) -> tuple[str, tuple[str, ...], int]:
|
||||
user_run: Final = tuple(
|
||||
itertools.takewhile(
|
||||
lambda message: message["role"] in _OLLAMA_USER_ROLES, (messages[i] for i in range(msg_i, len(messages)))
|
||||
)
|
||||
)
|
||||
user_parts: Final = tuple(
|
||||
_ollama_user_message_parts(model, message, msg_i + offset) for offset, message in enumerate(user_run)
|
||||
)
|
||||
user_content_str: Final = "\n".join("\n".join(texts) for texts, _ in user_parts if texts)
|
||||
image_urls: Final = tuple(itertools.chain.from_iterable(urls for _, urls in user_parts))
|
||||
return user_content_str, image_urls, len(user_run)
|
||||
|
||||
|
||||
def ollama_pt(
|
||||
model: str, messages: list
|
||||
) -> (
|
||||
str | OllamaVisionModelObject
|
||||
): # https://github.com/ollama/ollama/blob/af4cf55884ac54b9e637cd71dadfe9b7a5685877/docs/modelfile.md#template
|
||||
user_message_types: Final = {"user", "tool", "function"}
|
||||
msg_i = 0
|
||||
images: Final = []
|
||||
prompt = ""
|
||||
while msg_i < len(messages):
|
||||
init_msg_i = msg_i
|
||||
user_content_str = ""
|
||||
## MERGE CONSECUTIVE USER CONTENT ##
|
||||
while msg_i < len(messages) and messages[msg_i]["role"] in user_message_types:
|
||||
msg_content = messages[msg_i].get("content")
|
||||
if msg_content:
|
||||
if isinstance(msg_content, list):
|
||||
for m in msg_content:
|
||||
if m.get("type", "") == "image_url":
|
||||
if isinstance(m["image_url"], str):
|
||||
images.append(m["image_url"])
|
||||
elif isinstance(m["image_url"], dict):
|
||||
images.append(m["image_url"]["url"])
|
||||
elif m.get("type", "") == "text":
|
||||
user_content_str += m["text"]
|
||||
else:
|
||||
# Tool message content will always be a string
|
||||
user_content_str += msg_content
|
||||
|
||||
msg_i += 1
|
||||
user_content_str, image_urls, user_run_len = _ollama_user_turn(model, messages, msg_i)
|
||||
images.extend(image_urls)
|
||||
msg_i += user_run_len
|
||||
|
||||
if user_content_str:
|
||||
prompt += f"### User:\n{user_content_str}\n\n"
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ from litellm.types.utils import (
|
|||
ModelResponseStream,
|
||||
ProviderField,
|
||||
StreamingChoices,
|
||||
generate_id,
|
||||
)
|
||||
|
||||
from ..common_utils import OllamaError, OllamaModelInfo, convert_image
|
||||
|
|
@ -524,6 +525,7 @@ class OllamaConfig(BaseConfig):
|
|||
class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
|
||||
def __init__(self, streaming_response, sync_stream: bool, json_mode: bool | None = False):
|
||||
super().__init__(streaming_response, sync_stream, json_mode)
|
||||
self.response_id: Final[str] = generate_id()
|
||||
self.started_reasoning_content: bool = False
|
||||
self.finished_reasoning_content: bool = False
|
||||
self.streamed_content: bool = False
|
||||
|
|
@ -622,6 +624,7 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
|
|||
content = self._hold_json_object_start(text)
|
||||
|
||||
return ModelResponseStream(
|
||||
id=self.response_id,
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
|
|
@ -641,24 +644,26 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
|
|||
# Return reasoning content as ModelResponseStream so UIs can render it
|
||||
thinking_content: Final = chunk.get("thinking") or ""
|
||||
return ModelResponseStream(
|
||||
id=self.response_id,
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(reasoning_content=thinking_content),
|
||||
)
|
||||
]
|
||||
],
|
||||
)
|
||||
else:
|
||||
# In this case, 'thinking' is not present in the chunk, chunk["done"] is false,
|
||||
# and chunk["response"] is falsy (None or empty string),
|
||||
# but Ollama is just starting to stream, so it should be processed as a normal dict
|
||||
return ModelResponseStream(
|
||||
id=self.response_id,
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(reasoning_content=""),
|
||||
)
|
||||
]
|
||||
],
|
||||
)
|
||||
# raise Exception(f"Unable to parse ollama chunk - {chunk}")
|
||||
except Exception as e:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import time
|
||||
import uuid
|
||||
from collections.abc import Sequence
|
||||
from itertools import filterfalse
|
||||
from typing import Any, Final, cast
|
||||
|
||||
import litellm
|
||||
|
|
@ -81,6 +82,19 @@ def _delta_has_signed_thinking_block(delta: object) -> bool:
|
|||
return any(isinstance(b, dict) and (b.get("signature") or b.get("data")) for b in blocks)
|
||||
|
||||
|
||||
def _delta_carries_output(delta: ChatCompletionDelta) -> bool:
|
||||
fields: Final = (
|
||||
delta.content,
|
||||
getattr(delta, "reasoning_content", None),
|
||||
delta.tool_calls,
|
||||
delta.function_call,
|
||||
getattr(delta, "annotations", None),
|
||||
getattr(delta, "images", None),
|
||||
getattr(delta, "audio", None),
|
||||
)
|
||||
return any(fields) or _delta_has_signed_thinking_block(delta)
|
||||
|
||||
|
||||
class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
||||
"""
|
||||
Async iterator for processing streaming responses from the Responses API.
|
||||
|
|
@ -119,7 +133,8 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
self.completed_response = None
|
||||
self.final_text: str = ""
|
||||
self._cached_item_id: str | None = None
|
||||
self._message_output_index: int = 0
|
||||
self._message_output_index: int | None = None
|
||||
self._reasoning_output_index: int | None = None
|
||||
self._cached_response_id: str | None = None
|
||||
self._buffered_chunk: ModelResponseStream | None = None
|
||||
self._upstream_exhausted: bool = False
|
||||
|
|
@ -130,7 +145,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
self._tool_item_id_by_call_id: dict[str, str] = {} # mutable-ok: filled per call id as tool call events stream
|
||||
self._tool_call_id_by_index: dict[int, str] = {}
|
||||
self._ambiguous_tool_call_indexes: set[int] = set()
|
||||
self._next_tool_output_index: int = 1 # output_index=0 reserved for the message item
|
||||
self._next_output_index: int = 0
|
||||
self._final_tool_events_queued: bool = False
|
||||
self._sequence_number: int = 0
|
||||
self._cached_reasoning_item_id: str | None = None
|
||||
|
|
@ -155,11 +170,39 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
existing: Final = self._tool_output_index_by_call_id.get(call_id)
|
||||
if existing is not None:
|
||||
return existing
|
||||
idx: Final = self._next_tool_output_index
|
||||
self._next_tool_output_index += 1
|
||||
idx: Final = self._allocate_output_index()
|
||||
self._tool_output_index_by_call_id[call_id] = idx
|
||||
return idx
|
||||
|
||||
def _allocate_output_index(self) -> int:
|
||||
idx: Final = self._next_output_index
|
||||
self._next_output_index += 1
|
||||
return idx
|
||||
|
||||
def _message_index(self) -> int:
|
||||
if self._message_output_index is None:
|
||||
self._message_output_index = self._allocate_output_index()
|
||||
return self._message_output_index
|
||||
|
||||
def _reasoning_index(self) -> int:
|
||||
if self._reasoning_output_index is None:
|
||||
self._reasoning_output_index = self._allocate_output_index()
|
||||
return self._reasoning_output_index
|
||||
|
||||
def _streamed_output_index(self, item: object) -> int | None:
|
||||
match getattr(item, "type", None):
|
||||
case "message":
|
||||
return self._message_output_index
|
||||
case "reasoning":
|
||||
return self._reasoning_output_index
|
||||
case _:
|
||||
call_id: Final = getattr(item, "call_id", None) or str(getattr(item, "id", "")).removeprefix("ws_")
|
||||
return self._tool_output_index_by_call_id.get(str(call_id))
|
||||
|
||||
def _streamed_output_position(self, item: object) -> tuple[bool, int]:
|
||||
index: Final = self._streamed_output_index(item)
|
||||
return (index is None, index or 0)
|
||||
|
||||
def _normalize_tool_call_index(self, tool_call: object) -> int | None:
|
||||
idx_raw: Final = tool_call.get("index") if isinstance(tool_call, dict) else getattr(tool_call, "index", None)
|
||||
if idx_raw is None:
|
||||
|
|
@ -572,7 +615,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
self._sequence_number += 1
|
||||
event: Final = OutputItemAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
|
||||
output_index=self._message_output_index,
|
||||
output_index=self._message_index(),
|
||||
item=BaseLiteLLMOpenAIResponseObject(
|
||||
**{
|
||||
"id": self._cached_item_id,
|
||||
|
|
@ -594,7 +637,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
event: Final = ContentPartAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.CONTENT_PART_ADDED,
|
||||
item_id=self._cached_item_id,
|
||||
output_index=self._message_output_index,
|
||||
output_index=self._message_index(),
|
||||
content_index=0,
|
||||
part=BaseLiteLLMOpenAIResponseObject(**{"type": "output_text", "text": "", "annotations": []}),
|
||||
)
|
||||
|
|
@ -606,15 +649,10 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
self._cached_item_id = f"msg_{uuid.uuid4()}"
|
||||
self.sent_message_item_added_event = True
|
||||
self.sent_content_part_added_event = True
|
||||
if self._cached_reasoning_item_id is not None:
|
||||
self._message_output_index = self._next_tool_output_index
|
||||
self._next_tool_output_index += 1
|
||||
else:
|
||||
self._message_output_index = 0
|
||||
self._sequence_number += 1
|
||||
event: Final = OutputItemAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
|
||||
output_index=self._message_output_index,
|
||||
output_index=self._message_index(),
|
||||
item=BaseLiteLLMOpenAIResponseObject(
|
||||
**{
|
||||
"id": self._cached_item_id,
|
||||
|
|
@ -699,7 +737,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
return ReasoningSummaryTextDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DONE,
|
||||
item_id=reasoning_item_id,
|
||||
output_index=0,
|
||||
output_index=self._reasoning_index(),
|
||||
sequence_number=sequence_number,
|
||||
summary_index=0,
|
||||
text=reasoning_content,
|
||||
|
|
@ -730,7 +768,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
return ReasoningSummaryPartDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.REASONING_SUMMARY_PART_DONE,
|
||||
item_id=reasoning_item_id,
|
||||
output_index=0,
|
||||
output_index=self._reasoning_index(),
|
||||
sequence_number=sequence_number,
|
||||
summary_index=0,
|
||||
part=BaseLiteLLMOpenAIResponseObject(
|
||||
|
|
@ -748,7 +786,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
return OutputTextDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DONE,
|
||||
item_id=self._cached_item_id,
|
||||
output_index=self._message_output_index,
|
||||
output_index=self._message_index(),
|
||||
content_index=0,
|
||||
text=getattr(litellm_complete_object.choices[0].message, "content", "") or "",
|
||||
)
|
||||
|
|
@ -775,7 +813,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
return ContentPartDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.CONTENT_PART_DONE,
|
||||
item_id=self._cached_item_id,
|
||||
output_index=self._message_output_index,
|
||||
output_index=self._message_index(),
|
||||
content_index=0,
|
||||
part=part,
|
||||
)
|
||||
|
|
@ -794,7 +832,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
)
|
||||
return OutputItemDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE,
|
||||
output_index=self._message_output_index,
|
||||
output_index=self._message_index(),
|
||||
sequence_number=1,
|
||||
item=BaseLiteLLMOpenAIResponseObject(
|
||||
**{
|
||||
|
|
@ -841,7 +879,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
"""
|
||||
return OutputItemDoneEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE,
|
||||
output_index=0,
|
||||
output_index=self._reasoning_index(),
|
||||
sequence_number=sequence_number,
|
||||
item=BaseLiteLLMOpenAIResponseObject(
|
||||
**{
|
||||
|
|
@ -939,6 +977,8 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
if not chunk.choices:
|
||||
return
|
||||
delta: Final = chunk.choices[0].delta
|
||||
if chunk.choices[0].finish_reason is None and not _delta_carries_output(delta):
|
||||
return
|
||||
|
||||
self._sequence_number += 1
|
||||
self.sent_output_item_added_event = True
|
||||
|
|
@ -952,7 +992,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
|
||||
event = OutputItemAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED,
|
||||
output_index=0,
|
||||
output_index=self._reasoning_index(),
|
||||
item=BaseLiteLLMOpenAIResponseObject(
|
||||
**{
|
||||
"id": self._cached_reasoning_item_id,
|
||||
|
|
@ -1176,7 +1216,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
event = OutputTextAnnotationAddedEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_ANNOTATION_ADDED,
|
||||
item_id=item_id,
|
||||
output_index=self._message_output_index,
|
||||
output_index=self._message_index(),
|
||||
content_index=0,
|
||||
annotation_index=idx,
|
||||
annotation=annotation_dict,
|
||||
|
|
@ -1196,7 +1236,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
return ReasoningSummaryTextDeltaEvent(
|
||||
type=ResponsesAPIStreamEvents.REASONING_SUMMARY_TEXT_DELTA,
|
||||
item_id=self._cached_reasoning_item_id,
|
||||
output_index=0,
|
||||
output_index=self._reasoning_index(),
|
||||
delta=reasoning_content,
|
||||
)
|
||||
|
||||
|
|
@ -1209,7 +1249,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
text_delta_event: Final = OutputTextDeltaEvent(
|
||||
type=ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA,
|
||||
item_id=item_id,
|
||||
output_index=self._message_output_index,
|
||||
output_index=self._message_index(),
|
||||
content_index=0,
|
||||
delta=delta_content,
|
||||
)
|
||||
|
|
@ -1266,7 +1306,13 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
|
|||
"reasoning",
|
||||
self._cached_reasoning_item_id,
|
||||
)
|
||||
return reasoning_aligned
|
||||
streamed_items: Final = filterfalse(self._is_unstreamed_empty_message, reasoning_aligned)
|
||||
return tuple(sorted(streamed_items, key=self._streamed_output_position))
|
||||
|
||||
def _is_unstreamed_empty_message(self, item: object) -> bool:
|
||||
if getattr(item, "type", None) != "message" or self._message_output_index is not None:
|
||||
return False
|
||||
return not any(getattr(part, "text", None) for part in getattr(item, "content", None) or ())
|
||||
|
||||
def _emit_terminal_response_event(
|
||||
self, litellm_model_response: ModelResponse
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
import asyncio
|
||||
import itertools
|
||||
import json
|
||||
import uuid
|
||||
|
|
@ -8,10 +9,11 @@ from typing import Final
|
|||
import anthropic
|
||||
import openai
|
||||
import pytest
|
||||
from openai.types.chat import ChatCompletionChunk
|
||||
from openai.types import Completion
|
||||
from openai.types.chat import ChatCompletion, ChatCompletionChunk
|
||||
from openai.types.chat.chat_completion_chunk import Choice as ChunkChoice
|
||||
from openai.types.chat.chat_completion_chunk import ChoiceDeltaToolCall
|
||||
from integration._support.client import Gateway, eventually
|
||||
from integration._support.client import Gateway, eventually, object_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
|
@ -343,6 +345,7 @@ async def test_async_openai_sdk_stream_answers_the_tool_result_in_plain_text(gat
|
|||
extra_body={"cache": _NO_CACHE},
|
||||
)
|
||||
chunks: Final = [chunk async for chunk in stream]
|
||||
assert {chunk.id for chunk in chunks} == {chunks[0].id}
|
||||
choices: Final = tuple(_stream_choices(chunks))
|
||||
assert "".join(choice.delta.content or "" for choice in choices) == _ANSWER
|
||||
assert tuple(_delta_tool_calls(choices)) == ()
|
||||
|
|
@ -497,6 +500,7 @@ async def test_async_openai_sdk_responses_stream_emits_the_function_call_item(ga
|
|||
assert json.loads(item.arguments) == _ARGUMENTS
|
||||
completed: Final = [event for event in events if event.type == "response.completed"]
|
||||
assert len(completed) == 1
|
||||
assert [item.type for item in completed[0].response.output] == ["function_call"], completed[0].response.output
|
||||
final_calls: Final = [item for item in completed[0].response.output if item.type == "function_call"]
|
||||
assert [(item.name, json.loads(item.arguments)) for item in final_calls] == [("get_weather", _ARGUMENTS)]
|
||||
assert completed[0].response.output_text == ""
|
||||
|
|
@ -640,8 +644,8 @@ def test_ollama_server_error_on_the_tool_result_turn_does_not_take_the_deploymen
|
|||
pytest.param("r" * 5120, "r" * 5120, id="5kb-string-forwarded-intact"),
|
||||
pytest.param(
|
||||
[{"type": "text", "text": "Paris: 22 degrees"}, {"type": "text", "text": "clear skies"}],
|
||||
"Paris: 22 degreesclear skies",
|
||||
id="text-parts-joined",
|
||||
"Paris: 22 degrees\nclear skies",
|
||||
id="text-parts-joined-by-newline",
|
||||
),
|
||||
],
|
||||
)
|
||||
|
|
@ -668,6 +672,7 @@ def test_the_same_tool_result_twice_is_forwarded_twice_under_one_instruction(gat
|
|||
prompt: Final = _prompt_of(_only_generate(wire))
|
||||
_assert_instructed_once(prompt, "get_weather")
|
||||
assert prompt.count(_RESULT) == 2, prompt
|
||||
assert f"### User:\n{_RESULT}\n{_RESULT}\n\n" in prompt, prompt
|
||||
assert prompt.count("### User:") == 2 and prompt.count("### Assistant:") == 1, prompt
|
||||
|
||||
|
||||
|
|
@ -707,15 +712,16 @@ def test_a_non_function_json_answer_is_returned_as_text(gateway: Gateway) -> Non
|
|||
assert _spend_row(completion.id) == _billed(model)
|
||||
|
||||
|
||||
def test_int_tool_result_content_fails_in_the_response_body_and_leaves_the_deployment_serving(gateway: Gateway) -> None:
|
||||
def test_int_tool_result_content_is_a_400_naming_the_field_and_leaves_the_deployment_serving(gateway: Gateway) -> None:
|
||||
with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
code, text = _post(
|
||||
gateway, "/v1/chat/completions", {"model": model, "messages": _second_turn(22), "tools": [_WEATHER_TOOL]}
|
||||
)
|
||||
assert code >= 400, text
|
||||
assert code == 400, text
|
||||
error: Final = _JSON_OBJECT.validate_json(text)["error"]
|
||||
assert isinstance(error, dict) and isinstance(error["message"], str) and error["message"], text
|
||||
assert isinstance(error, dict) and isinstance(error["message"], str), text
|
||||
assert "content" in error["message"] and "tool message" in error["message"], text
|
||||
assert _generate_calls(wire) == ()
|
||||
payload: Final = _post_chat(gateway, model, _second_turn(), tools=[_WEATHER_TOOL])
|
||||
assert json.dumps(payload["choices"]).count(_ANSWER) == 1, payload
|
||||
|
|
@ -794,3 +800,601 @@ def _deployments(gateway: Gateway) -> list[JsonValue]:
|
|||
data: Final = gateway.get("/model/info")["data"]
|
||||
assert isinstance(data, list)
|
||||
return data
|
||||
|
||||
|
||||
_FOLLOW_UP: Final = "Is it windy there too?"
|
||||
_TIME_CALL_ID: Final = "call_prompt_tools_2"
|
||||
_TIME_RESULT: Final = "Paris: 14:05 local time"
|
||||
_THINKING: Final = "The tool result already answers the question, so reply in plain text."
|
||||
_ANTHROPIC_TIME_TOOL: Final[dict[str, JsonValue]] = {
|
||||
"name": "get_time",
|
||||
"description": "Local time for a city",
|
||||
"input_schema": _PARAMETERS,
|
||||
}
|
||||
_RESPONSES_TIME_TOOL: Final[dict[str, JsonValue]] = {
|
||||
"type": "function",
|
||||
"name": "get_time",
|
||||
"description": "Local time for a city",
|
||||
"parameters": _PARAMETERS,
|
||||
}
|
||||
_AUDIO_PART: Final[dict[str, JsonValue]] = {"type": "input_audio", "input_audio": {"data": "UklGRg==", "format": "wav"}}
|
||||
|
||||
|
||||
def _ndjson_line(**fields: JsonValue) -> bytes:
|
||||
return json.dumps({"model": _BACKEND, "created_at": "2026-10-07T00:00:00Z", **fields}).encode() + b"\n"
|
||||
|
||||
|
||||
def _ndjson_reply(lines: Sequence[bytes]) -> Reply:
|
||||
final: Final = _ndjson_line(
|
||||
response="", done=True, done_reason="stop", prompt_eval_count=30, eval_count=12
|
||||
)
|
||||
return Reply(content_type="application/x-ndjson", chunks=(*lines, final))
|
||||
|
||||
|
||||
def _thinking_reply(thinking: str) -> Reply:
|
||||
pieces: Final = tuple(thinking[index : index + 7] for index in range(0, len(thinking), 7))
|
||||
return _ndjson_reply(
|
||||
(
|
||||
_ndjson_line(response="", done=False),
|
||||
*(_ndjson_line(response="", thinking=piece, done=False) for piece in pieces),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _completion_texts(chunks: Sequence[Completion]) -> Iterator[str]:
|
||||
for chunk in chunks:
|
||||
yield from (choice.text or "" for choice in chunk.choices)
|
||||
|
||||
|
||||
def _raw_stream_frames(gateway: Gateway, path: str, body: dict[str, JsonValue]) -> tuple[str, ...]:
|
||||
with gateway.client.stream(
|
||||
"POST", path, json={**body, "cache": _NO_CACHE}, headers={"Authorization": f"Bearer {gateway.key}"}
|
||||
) as response:
|
||||
assert response.status_code == 200, response.read()
|
||||
return tuple(line.removeprefix("data: ") for line in response.iter_lines() if line.startswith("data: "))
|
||||
|
||||
|
||||
def _raw_choice_contents(choices: JsonValue) -> Iterator[str]:
|
||||
assert isinstance(choices, list)
|
||||
for choice in choices:
|
||||
content: Final = object_value(object_value(choice)["delta"]).get("content")
|
||||
if isinstance(content, str):
|
||||
yield content
|
||||
|
||||
|
||||
def _raw_delta_contents(payloads: Sequence[dict[str, JsonValue]]) -> Iterator[str]:
|
||||
for payload in payloads:
|
||||
yield from _raw_choice_contents(payload["choices"])
|
||||
|
||||
|
||||
def _cache_rows(model: str) -> list[dict[str, JsonValue]]:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT request_id, cache_hit, status FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)
|
||||
),
|
||||
lambda found: len(found) >= 2,
|
||||
seconds=70,
|
||||
)
|
||||
assert len(rows) == 2, rows
|
||||
return rows
|
||||
|
||||
|
||||
def _two_call_turn() -> list[dict[str, JsonValue]]:
|
||||
return [
|
||||
{"role": "user", "content": _QUESTION},
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "tool_use", "id": _CALL_ID, "name": "get_weather", "input": _ARGUMENTS},
|
||||
{"type": "tool_use", "id": _TIME_CALL_ID, "name": "get_time", "input": _ARGUMENTS},
|
||||
],
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "tool_result", "tool_use_id": _CALL_ID, "content": _RESULT},
|
||||
{"type": "tool_result", "tool_use_id": _TIME_CALL_ID, "content": _TIME_RESULT},
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def _anthropic_result_with_text() -> list[dict[str, JsonValue]]:
|
||||
return [
|
||||
*_anthropic_second_turn()[:2],
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "tool_result", "tool_use_id": _CALL_ID, "content": _RESULT},
|
||||
{"type": "text", "text": _FOLLOW_UP},
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def _anthropic_int_result() -> list[dict[str, JsonValue]]:
|
||||
return [
|
||||
*_anthropic_second_turn()[:2],
|
||||
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": _CALL_ID, "content": 22}]},
|
||||
]
|
||||
|
||||
|
||||
def _responses_two_outputs() -> list[dict[str, JsonValue]]:
|
||||
return [
|
||||
{"role": "user", "content": _QUESTION},
|
||||
{"type": "function_call", "call_id": _CALL_ID, "name": "get_weather", "arguments": json.dumps(_ARGUMENTS)},
|
||||
{"type": "function_call", "call_id": _TIME_CALL_ID, "name": "get_time", "arguments": json.dumps(_ARGUMENTS)},
|
||||
{"type": "function_call_output", "call_id": _CALL_ID, "output": _RESULT},
|
||||
{"type": "function_call_output", "call_id": _TIME_CALL_ID, "output": _TIME_RESULT},
|
||||
]
|
||||
|
||||
|
||||
def _responses_int_output() -> list[dict[str, JsonValue]]:
|
||||
return [*_responses_second_turn()[:2], {"type": "function_call_output", "call_id": _CALL_ID, "output": 22}]
|
||||
|
||||
|
||||
def _error_message(text: str) -> str:
|
||||
error: Final = _JSON_OBJECT.validate_json(text)["error"]
|
||||
assert isinstance(error, dict) and isinstance(error["message"], str), text
|
||||
return error["message"]
|
||||
|
||||
|
||||
async def _chat_attempt(client: openai.AsyncOpenAI, model: str, messages: list[dict[str, JsonValue]]) -> ChatCompletion:
|
||||
return await client.chat.completions.create(
|
||||
model=model,
|
||||
messages=messages, # pyright: ignore[reportArgumentType] # plain JSON messages
|
||||
tools=[_WEATHER_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool
|
||||
extra_body={"cache": _NO_CACHE},
|
||||
)
|
||||
|
||||
|
||||
def test_openai_sdk_completions_stream_keeps_one_id(gateway: Gateway) -> None:
|
||||
with _ollama_server(lambda _: _streamed_reply(_ANSWER)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
stream: Final = _openai_client(gateway).completions.create(
|
||||
model=model, prompt=_QUESTION, stream=True, extra_body={"cache": _NO_CACHE}
|
||||
)
|
||||
chunks: Final = list(stream)
|
||||
assert {chunk.id for chunk in chunks} == {chunks[0].id}, [chunk.id for chunk in chunks]
|
||||
assert "".join(_completion_texts(chunks)) == _ANSWER
|
||||
body: Final = _only_generate(wire)
|
||||
assert body["stream"] is True
|
||||
assert body["model"] == _BACKEND
|
||||
assert body["prompt"] == _QUESTION, body
|
||||
assert _spend_row(chunks[0].id) == _billed(model)
|
||||
|
||||
|
||||
def test_raw_chat_stream_frames_share_one_id_and_end_with_done(gateway: Gateway) -> None:
|
||||
with _ollama_server(lambda _: _streamed_reply(_ANSWER)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
frames: Final = _raw_stream_frames(
|
||||
gateway,
|
||||
"/v1/chat/completions",
|
||||
{"model": model, "messages": _second_turn(), "tools": [_WEATHER_TOOL], "stream": True},
|
||||
)
|
||||
assert frames[-1] == "[DONE]", frames
|
||||
payloads: Final = [_JSON_OBJECT.validate_json(frame) for frame in frames[:-1]]
|
||||
assert {payload["id"] for payload in payloads} == {payloads[0]["id"]}, [payload["id"] for payload in payloads]
|
||||
assert "".join(_raw_delta_contents(payloads)) == _ANSWER
|
||||
body: Final = _only_generate(wire)
|
||||
assert body["stream"] is True
|
||||
_assert_tool_turn(_prompt_of(body))
|
||||
identity: Final = payloads[0]["id"]
|
||||
assert isinstance(identity, str)
|
||||
assert _spend_row(identity) == _billed(model)
|
||||
|
||||
|
||||
async def test_cached_stream_replay_keeps_one_id_and_reaches_ollama_once(gateway: Gateway) -> None:
|
||||
messages: Final[list[dict[str, JsonValue]]] = [{"role": "user", "content": f"{_QUESTION} ({uuid.uuid4().hex})"}]
|
||||
with _ollama_server(lambda _: _streamed_reply(_ANSWER)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
client: Final = _async_openai_client(gateway)
|
||||
first: Final = [
|
||||
chunk
|
||||
async for chunk in await client.chat.completions.create(
|
||||
model=model,
|
||||
messages=messages, # pyright: ignore[reportArgumentType] # plain JSON messages
|
||||
stream=True,
|
||||
)
|
||||
]
|
||||
assert {chunk.id for chunk in first} == {first[0].id}, [chunk.id for chunk in first]
|
||||
assert "".join(choice.delta.content or "" for choice in _stream_choices(first)) == _ANSWER
|
||||
assert _spend_row(first[0].id) == _billed(model)
|
||||
second: Final = [
|
||||
chunk
|
||||
async for chunk in await client.chat.completions.create(
|
||||
model=model,
|
||||
messages=messages, # pyright: ignore[reportArgumentType] # plain JSON messages
|
||||
stream=True,
|
||||
)
|
||||
]
|
||||
assert {chunk.id for chunk in second} == {second[0].id}, [chunk.id for chunk in second]
|
||||
assert "".join(choice.delta.content or "" for choice in _stream_choices(second)) == _ANSWER
|
||||
assert len(_generate_calls(wire)) == 1
|
||||
rows: Final = _cache_rows(model)
|
||||
hits: Final = [row for row in rows if row["cache_hit"] == "True"]
|
||||
assert len(hits) == 1, rows
|
||||
hit_id: Final = hits[0]["request_id"]
|
||||
assert isinstance(hit_id, str) and hit_id.startswith(second[0].id), rows
|
||||
assert [row["request_id"] for row in rows if row["cache_hit"] != "True"] == [first[0].id], rows
|
||||
assert {row["status"] for row in rows} == {"success"}, rows
|
||||
|
||||
|
||||
def test_consecutive_user_messages_are_joined_by_a_newline(gateway: Gateway) -> None:
|
||||
with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
messages: Final[list[dict[str, JsonValue]]] = [*_first_turn(), {"role": "user", "content": _FOLLOW_UP}]
|
||||
_post_chat(gateway, model, messages, tools=[_WEATHER_TOOL])
|
||||
prompt: Final = _prompt_of(_only_generate(wire))
|
||||
_assert_instructed_once(prompt, "get_weather")
|
||||
assert f"### User:\n{_QUESTION}\n{_FOLLOW_UP}\n\n" in prompt, prompt
|
||||
assert prompt.count("### User:") == 1, prompt
|
||||
|
||||
|
||||
def test_a_tool_result_followed_by_user_text_is_joined_by_a_newline(gateway: Gateway) -> None:
|
||||
with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
messages: Final[list[dict[str, JsonValue]]] = [*_second_turn(), {"role": "user", "content": _FOLLOW_UP}]
|
||||
payload: Final = _post_chat(gateway, model, messages, tools=[_WEATHER_TOOL])
|
||||
assert json.dumps(payload["choices"]).count(_ANSWER) == 1, payload
|
||||
prompt: Final = _prompt_of(_only_generate(wire))
|
||||
_assert_instructed_once(prompt, "get_weather")
|
||||
assert (
|
||||
f"### User:\n{_QUESTION}\n\n### Assistant:\n{_CALL_JSON}\n\n### User:\n{_RESULT}\n{_FOLLOW_UP}\n\n" in prompt
|
||||
), prompt
|
||||
assert prompt.count("### User:") == 2, prompt
|
||||
|
||||
|
||||
def test_anthropic_sdk_two_tool_results_in_one_turn_are_separated(gateway: Gateway) -> None:
|
||||
with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
message: Final = _anthropic_client(gateway).messages.create(
|
||||
model=model,
|
||||
max_tokens=64,
|
||||
messages=_two_call_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages
|
||||
tools=[_ANTHROPIC_TOOL, _ANTHROPIC_TIME_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tools
|
||||
extra_body={"cache": _NO_CACHE},
|
||||
)
|
||||
assert message.stop_reason == "end_turn"
|
||||
assert [(block.type, getattr(block, "text", None)) for block in message.content] == [("text", _ANSWER)]
|
||||
prompt: Final = _prompt_of(_only_generate(wire))
|
||||
_assert_instructed_once(prompt, "get_weather", "get_time")
|
||||
assert f"### User:\n{_RESULT}\n{_TIME_RESULT}\n\n" in prompt, prompt
|
||||
assert prompt.count("### User:") == 2 and prompt.count("### Assistant:") == 1, prompt
|
||||
assistant: Final = prompt.split("### Assistant:\n", 1)[1].split("### User:", 1)[0]
|
||||
assert "get_weather" in assistant and "get_time" in assistant, assistant
|
||||
assert _spend_row(message.id) == _billed(model)
|
||||
|
||||
|
||||
async def test_async_anthropic_sdk_stream_separates_a_tool_result_from_user_text_in_one_turn(
|
||||
gateway: Gateway,
|
||||
) -> None:
|
||||
with _ollama_server(lambda _: _streamed_reply(_ANSWER)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
async with _async_anthropic_client(gateway).messages.stream(
|
||||
model=model,
|
||||
max_tokens=64,
|
||||
messages=_anthropic_result_with_text(), # pyright: ignore[reportArgumentType] # plain JSON messages
|
||||
tools=[_ANTHROPIC_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool
|
||||
extra_body={"cache": _NO_CACHE},
|
||||
) as stream:
|
||||
texts: Final = [event.text async for event in stream if event.type == "text"]
|
||||
final: Final = await stream.get_final_message()
|
||||
assert "".join(texts) == _ANSWER
|
||||
assert final.stop_reason == "end_turn"
|
||||
body: Final = _only_generate(wire)
|
||||
assert body["stream"] is True
|
||||
prompt: Final = _prompt_of(body)
|
||||
_assert_instructed_once(prompt, "get_weather")
|
||||
assert (
|
||||
f"### User:\n{_QUESTION}\n\n### Assistant:\n{_CALL_JSON}\n\n### User:\n{_RESULT}\n{_FOLLOW_UP}\n\n" in prompt
|
||||
), prompt
|
||||
assert _spend_row(final.id) == _billed(model)
|
||||
|
||||
|
||||
def test_anthropic_sdk_int_tool_result_content_is_dropped_before_the_prompt(gateway: Gateway) -> None:
|
||||
with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
message: Final = _anthropic_client(gateway).messages.create(
|
||||
model=model,
|
||||
max_tokens=64,
|
||||
messages=_anthropic_int_result(), # pyright: ignore[reportArgumentType] # plain JSON messages
|
||||
tools=[_ANTHROPIC_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool
|
||||
extra_body={"cache": _NO_CACHE},
|
||||
)
|
||||
assert [(block.type, getattr(block, "text", None)) for block in message.content] == [("text", _ANSWER)]
|
||||
prompt: Final = _prompt_of(_only_generate(wire))
|
||||
_assert_instructed_once(prompt, "get_weather")
|
||||
assert f"### User:\n{_QUESTION}\n\n### Assistant:\n{_CALL_JSON}\n\n### System:\n" in prompt, prompt
|
||||
assert "22" not in prompt.split("### Assistant:\n", 1)[1], prompt
|
||||
assert _spend_row(message.id) == _billed(model)
|
||||
|
||||
|
||||
def test_openai_sdk_responses_two_function_outputs_are_separated(gateway: Gateway) -> None:
|
||||
with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
response: Final = _openai_client(gateway).responses.create(
|
||||
model=model,
|
||||
input=_responses_two_outputs(), # pyright: ignore[reportArgumentType] # plain JSON input items
|
||||
tools=[_RESPONSES_TOOL, _RESPONSES_TIME_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tools
|
||||
store=False,
|
||||
extra_body={"cache": _NO_CACHE},
|
||||
)
|
||||
assert [item.type for item in response.output] == ["message"]
|
||||
assert response.output_text == _ANSWER
|
||||
prompt: Final = _prompt_of(_only_generate(wire))
|
||||
_assert_instructed_once(prompt, "get_weather", "get_time")
|
||||
assert f"### User:\n{_RESULT}\n{_TIME_RESULT}\n\n" in prompt, prompt
|
||||
assert prompt.count("### User:") == 2 and prompt.count("### Assistant:") == 1, prompt
|
||||
assert _model_spend_rows(model, 1)[0]["status"] == "success"
|
||||
|
||||
|
||||
def test_openai_sdk_responses_int_function_output_reaches_the_prompt_as_text(gateway: Gateway) -> None:
|
||||
with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
response: Final = _openai_client(gateway).responses.create(
|
||||
model=model,
|
||||
input=_responses_int_output(), # pyright: ignore[reportArgumentType] # plain JSON input items
|
||||
tools=[_RESPONSES_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool
|
||||
store=False,
|
||||
extra_body={"cache": _NO_CACHE},
|
||||
)
|
||||
assert response.output_text == _ANSWER
|
||||
_assert_tool_turn(_prompt_of(_only_generate(wire)), "22")
|
||||
assert _model_spend_rows(model, 1)[0]["status"] == "success"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("messages", "message_ref", "detail"),
|
||||
[
|
||||
pytest.param(_second_turn(22), "tool message at index 2", "has int content", id="tool-int-content"),
|
||||
pytest.param(
|
||||
_second_turn({"text": _RESULT}), "tool message at index 2", "has dict content", id="tool-dict-content"
|
||||
),
|
||||
pytest.param(
|
||||
_second_turn([_RESULT]), "tool message at index 2", "has a str content part", id="tool-string-part"
|
||||
),
|
||||
pytest.param(
|
||||
_second_turn([{"type": "text", "text": 22}]),
|
||||
"tool message at index 2",
|
||||
"has a int text part",
|
||||
id="tool-int-text-part",
|
||||
),
|
||||
pytest.param(
|
||||
_second_turn([{"type": "text"}]),
|
||||
"tool message at index 2",
|
||||
"has a text part with no text",
|
||||
id="tool-text-part-without-text",
|
||||
),
|
||||
pytest.param(
|
||||
[{"role": "user", "content": [{"type": "text", "text": _QUESTION}, {"type": "image_url", "image_url": 22}]}],
|
||||
"user message at index 0",
|
||||
"has a int image_url",
|
||||
id="user-int-image-url",
|
||||
),
|
||||
pytest.param(
|
||||
[{"role": "user", "content": [{"type": "image_url", "image_url": {"detail": "high"}}]}],
|
||||
"user message at index 0",
|
||||
"has an image_url object without a url string",
|
||||
id="user-image-url-object-without-url",
|
||||
),
|
||||
pytest.param(
|
||||
[{"role": "user", "content": [{"type": "image_url"}]}],
|
||||
"user message at index 0",
|
||||
"has an image_url part with no image_url",
|
||||
id="user-image-url-part-without-image-url",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_malformed_content_is_a_400_naming_the_message_and_field(
|
||||
gateway: Gateway, messages: list[dict[str, JsonValue]], message_ref: str, detail: str
|
||||
) -> None:
|
||||
with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
code, text = _post(gateway, "/v1/chat/completions", {"model": model, "messages": messages, "tools": [_WEATHER_TOOL]})
|
||||
assert code == 400, text
|
||||
assert f"the {message_ref} {detail}" in _error_message(text), text
|
||||
assert _generate_calls(wire) == ()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"content",
|
||||
[
|
||||
pytest.param(22, id="user-int-content"),
|
||||
pytest.param({"text": _QUESTION}, id="user-dict-content"),
|
||||
pytest.param(["just a string"], id="user-string-part"),
|
||||
],
|
||||
)
|
||||
def test_non_list_user_content_is_rejected_before_it_reaches_ollama(gateway: Gateway, content: JsonValue) -> None:
|
||||
with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
code, text = _post(
|
||||
gateway,
|
||||
"/v1/chat/completions",
|
||||
{"model": model, "messages": [{"role": "user", "content": content}], "tools": [_WEATHER_TOOL]},
|
||||
)
|
||||
assert code >= 400, text
|
||||
assert _error_message(text), text
|
||||
assert _generate_calls(wire) == ()
|
||||
payload: Final = _post_chat(gateway, model, _second_turn(), tools=[_WEATHER_TOOL])
|
||||
assert json.dumps(payload["choices"]).count(_ANSWER) == 1, payload
|
||||
_assert_tool_turn(_prompt_of(_only_generate(wire)))
|
||||
|
||||
|
||||
def test_malformed_content_400_writes_a_failure_spend_row_and_leaves_the_deployment_serving(gateway: Gateway) -> None:
|
||||
with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
code, text = _post(
|
||||
gateway, "/v1/chat/completions", {"model": model, "messages": _second_turn(22), "tools": [_WEATHER_TOOL]}
|
||||
)
|
||||
assert code == 400, text
|
||||
assert "the tool message at index 2 has int content" in _error_message(text), text
|
||||
assert [row["status"] for row in _model_spend_rows(model, 1)] == ["failure"]
|
||||
assert _generate_calls(wire) == ()
|
||||
payload: Final = _post_chat(gateway, model, _second_turn(), tools=[_WEATHER_TOOL])
|
||||
assert json.dumps(payload["choices"]).count(_ANSWER) == 1, payload
|
||||
_assert_tool_turn(_prompt_of(_only_generate(wire)))
|
||||
assert sorted(str(row["status"]) for row in _model_spend_rows(model, 2)) == ["failure", "success"]
|
||||
|
||||
|
||||
def test_malformed_content_on_an_unauthenticated_request_is_a_401_before_translation(gateway: Gateway) -> None:
|
||||
with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
code, text = _post(
|
||||
gateway,
|
||||
"/v1/chat/completions",
|
||||
{"model": model, "messages": _second_turn(22), "tools": [_WEATHER_TOOL]},
|
||||
key=f"sk-not-a-key-{uuid.uuid4().hex}",
|
||||
)
|
||||
assert code == 401, text
|
||||
assert _generate_calls(wire) == ()
|
||||
|
||||
|
||||
def test_a_content_part_of_an_unknown_type_is_dropped_from_the_prompt(gateway: Gateway) -> None:
|
||||
with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
messages: Final[list[dict[str, JsonValue]]] = [
|
||||
{"role": "user", "content": [{"type": "text", "text": _QUESTION}, _AUDIO_PART]}
|
||||
]
|
||||
_post_chat(gateway, model, messages, tools=[_WEATHER_TOOL])
|
||||
prompt: Final = _prompt_of(_only_generate(wire))
|
||||
_assert_instructed_once(prompt, "get_weather")
|
||||
assert f"### User:\n{_QUESTION}\n\n" in prompt, prompt
|
||||
assert "UklGRg==" not in prompt and "input_audio" not in prompt, prompt
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"messages",
|
||||
[
|
||||
pytest.param(
|
||||
[
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "text", "text": ""},
|
||||
{"type": "text", "text": _QUESTION},
|
||||
{"type": "text", "text": ""},
|
||||
],
|
||||
}
|
||||
],
|
||||
id="empty-text-parts-dropped",
|
||||
),
|
||||
pytest.param([{"role": "user", "content": None}, *_first_turn()], id="null-content-skipped"),
|
||||
pytest.param([{"role": "user", "content": ""}, *_first_turn()], id="empty-string-skipped"),
|
||||
],
|
||||
)
|
||||
def test_empty_and_null_user_content_are_dropped_from_the_user_section(
|
||||
gateway: Gateway, messages: list[dict[str, JsonValue]]
|
||||
) -> None:
|
||||
with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
_post_chat(gateway, model, messages, tools=[_WEATHER_TOOL])
|
||||
prompt: Final = _prompt_of(_only_generate(wire))
|
||||
_assert_instructed_once(prompt, "get_weather")
|
||||
assert f"### User:\n{_QUESTION}\n\n" in prompt, prompt
|
||||
assert prompt.count("### User:") == 1, prompt
|
||||
|
||||
|
||||
async def test_async_openai_sdk_responses_stream_of_thinking_only_emits_one_reasoning_item(gateway: Gateway) -> None:
|
||||
with _ollama_server(lambda _: _thinking_reply(_THINKING)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
stream: Final = await _async_openai_client(gateway).responses.create(
|
||||
model=model,
|
||||
input=_responses_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON input items
|
||||
tools=[_RESPONSES_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool
|
||||
store=False,
|
||||
stream=True,
|
||||
extra_body={"cache": _NO_CACHE},
|
||||
)
|
||||
events: Final = [event async for event in stream]
|
||||
added: Final = [
|
||||
(event.item.type, event.output_index) for event in events if event.type == "response.output_item.added"
|
||||
]
|
||||
assert added == [("reasoning", 0)], [event.type for event in events]
|
||||
completed: Final = [event for event in events if event.type == "response.completed"]
|
||||
assert len(completed) == 1, [event.type for event in events]
|
||||
assert [item.type for item in completed[0].response.output] == ["reasoning"], completed[0].response.output
|
||||
assert completed[0].response.output_text == ""
|
||||
body: Final = _only_generate(wire)
|
||||
assert body["stream"] is True
|
||||
_assert_tool_turn(_prompt_of(body))
|
||||
assert _model_spend_rows(model, 1)[0]["status"] == "success"
|
||||
|
||||
|
||||
async def test_async_openai_sdk_responses_stream_with_an_empty_answer_keeps_one_message_item(gateway: Gateway) -> None:
|
||||
with _ollama_server(lambda _: _ndjson_reply(())) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
stream: Final = await _async_openai_client(gateway).responses.create(
|
||||
model=model,
|
||||
input=_responses_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON input items
|
||||
tools=[_RESPONSES_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool
|
||||
store=False,
|
||||
stream=True,
|
||||
extra_body={"cache": _NO_CACHE},
|
||||
)
|
||||
events: Final = [event async for event in stream]
|
||||
added: Final = [
|
||||
(event.item.type, event.output_index) for event in events if event.type == "response.output_item.added"
|
||||
]
|
||||
assert added == [("message", 0)], [event.type for event in events]
|
||||
completed: Final = [event for event in events if event.type == "response.completed"]
|
||||
assert len(completed) == 1, [event.type for event in events]
|
||||
assert [item.type for item in completed[0].response.output] == ["message"], completed[0].response.output
|
||||
assert completed[0].response.output_text == ""
|
||||
body: Final = _only_generate(wire)
|
||||
assert body["stream"] is True
|
||||
_assert_tool_turn(_prompt_of(body))
|
||||
assert _model_spend_rows(model, 1)[0]["status"] == "success"
|
||||
|
||||
|
||||
def test_openai_sdk_responses_stream_context_manager_gets_the_function_call_without_an_empty_message(
|
||||
gateway: Gateway,
|
||||
) -> None:
|
||||
with _ollama_server(lambda _: _streamed_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
with _openai_client(gateway).responses.stream(
|
||||
model=model,
|
||||
input=_QUESTION,
|
||||
tools=[_RESPONSES_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool
|
||||
store=False,
|
||||
extra_body={"cache": _NO_CACHE},
|
||||
) as stream:
|
||||
events: Final = list(stream)
|
||||
final: Final = stream.get_final_response()
|
||||
added: Final = [
|
||||
(event.item.type, event.output_index) for event in events if event.type == "response.output_item.added"
|
||||
]
|
||||
assert added == [("function_call", 0)], [event.type for event in events]
|
||||
assert [item.type for item in final.output] == ["function_call"], final.output
|
||||
call: Final = final.output[0]
|
||||
assert call.type == "function_call" and call.name == "get_weather"
|
||||
assert json.loads(call.arguments) == _ARGUMENTS
|
||||
assert final.output_text == ""
|
||||
body: Final = _only_generate(wire)
|
||||
assert body["stream"] is True
|
||||
_assert_instructed_once(_prompt_of(body), "get_weather")
|
||||
assert _model_spend_rows(model, 1)[0]["status"] == "success"
|
||||
|
||||
|
||||
async def test_concurrent_valid_and_malformed_turns_are_each_answered_in_their_own_shape(gateway: Gateway) -> None:
|
||||
with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
client: Final = _async_openai_client(gateway)
|
||||
results: Final = await asyncio.gather(
|
||||
*(_chat_attempt(client, model, _second_turn(22 if index % 3 == 0 else _RESULT)) for index in range(12)),
|
||||
return_exceptions=True,
|
||||
)
|
||||
answers: Final = [result for result in results if isinstance(result, ChatCompletion)]
|
||||
rejections: Final = [result for result in results if isinstance(result, openai.BadRequestError)]
|
||||
assert (len(answers), len(rejections)) == (8, 4), results
|
||||
assert {answer.choices[0].message.content for answer in answers} == {_ANSWER}
|
||||
for rejection in rejections:
|
||||
assert "the tool message at index 2 has int content" in str(rejection), rejection
|
||||
prompts: Final = [_prompt_of(_JSON_OBJECT.validate_json(request.body)) for request in _generate_calls(wire)]
|
||||
assert len(prompts) == 8, prompts
|
||||
for prompt in prompts:
|
||||
_assert_tool_turn(prompt)
|
||||
statuses: Final = sorted(str(row["status"]) for row in _model_spend_rows(model, 12))
|
||||
assert statuses == ["failure"] * 4 + ["success"] * 8, statuses
|
||||
for answer in answers:
|
||||
assert _spend_row(answer.id) == _billed(model)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,248 @@
|
|||
import json
|
||||
import uuid
|
||||
from collections.abc import Callable, Iterator, Sequence
|
||||
from contextlib import contextmanager
|
||||
from typing import Final
|
||||
|
||||
import openai
|
||||
from integration._support.client import Gateway, Scenario, eventually, object_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.wire import Reply, Request, Wire, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
_BACKEND: Final = "qwen3-bridge-items"
|
||||
_API_KEY: Final = "synthetic-hosted-vllm-key"
|
||||
_QUESTION: Final = "What is the weather in Paris?"
|
||||
_ANSWER: Final = "Paris is 22 degrees Celsius with clear skies."
|
||||
_PREFACE: Final = "Checking the weather."
|
||||
_REASONING: Final = "The user asks for the weather, so the weather tool applies."
|
||||
_CALL_ID: Final = "call_bridge_items_1"
|
||||
_ARGUMENTS: Final[dict[str, JsonValue]] = {"city": "Paris"}
|
||||
_TOOL: Final[dict[str, JsonValue]] = {
|
||||
"type": "function",
|
||||
"name": "get_weather",
|
||||
"description": "Weather for a city",
|
||||
"parameters": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]},
|
||||
}
|
||||
_NO_CACHE: Final[dict[str, JsonValue]] = {"no-cache": True}
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
_DISCOVERY_PROBE: Final = ("GET", "/v1/models")
|
||||
|
||||
|
||||
def _frame(identity: str, delta: dict[str, JsonValue], finish_reason: str | None = None) -> bytes:
|
||||
usage: Final = {"usage": {"prompt_tokens": 30, "completion_tokens": 12, "total_tokens": 42}} if finish_reason else {}
|
||||
body: Final = {
|
||||
"id": identity,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 1,
|
||||
"model": _BACKEND,
|
||||
"choices": [{"index": 0, "delta": delta, "finish_reason": finish_reason}],
|
||||
**usage,
|
||||
}
|
||||
return b"data: " + json.dumps(body).encode() + b"\n\n"
|
||||
|
||||
|
||||
def _sse_reply(identity: str, deltas: Sequence[dict[str, JsonValue]], finish_reason: str) -> Reply:
|
||||
return Reply(
|
||||
content_type="text/event-stream",
|
||||
chunks=(
|
||||
_frame(identity, {"role": "assistant", "content": ""}),
|
||||
*(_frame(identity, delta) for delta in deltas),
|
||||
_frame(identity, {}, finish_reason),
|
||||
b"data: [DONE]\n\n",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _tool_call_delta() -> dict[str, JsonValue]:
|
||||
return {
|
||||
"tool_calls": [
|
||||
{
|
||||
"index": 0,
|
||||
"id": _CALL_ID,
|
||||
"type": "function",
|
||||
"function": {"name": "get_weather", "arguments": json.dumps(_ARGUMENTS)},
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
def _is_discovery_probe(request: Request) -> bool:
|
||||
return (request.method, request.target) == _DISCOVERY_PROBE
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _vllm_server(respond: Callable[[Request], Reply]) -> Iterator[Wire]:
|
||||
with wire_server(
|
||||
lambda request: Reply(body=b'{"object":"list","data":[]}') if _is_discovery_probe(request) else respond(request)
|
||||
) as wire:
|
||||
yield wire
|
||||
|
||||
|
||||
def _bridged_vllm_model(scenario: Scenario, wire: Wire) -> str:
|
||||
return scenario.model(
|
||||
model=f"hosted_vllm/{_BACKEND}", api_base=wire.url + "/v1", api_key=_API_KEY, use_chat_completions_api=True
|
||||
)
|
||||
|
||||
|
||||
def _only_streamed_chat(wire: Wire, *tool_names: str) -> dict[str, JsonValue]:
|
||||
received: Final = tuple(request for request in wire.drain() if not _is_discovery_probe(request))
|
||||
assert [(request.method, request.target) for request in received] == [("POST", "/v1/chat/completions")]
|
||||
assert received[0].headers["authorization"] == f"Bearer {_API_KEY}"
|
||||
body: Final = _JSON_OBJECT.validate_json(received[0].body)
|
||||
assert body["stream"] is True
|
||||
tools: Final = body.get("tools", [])
|
||||
assert isinstance(tools, list)
|
||||
assert [object_value(object_value(tool)["function"])["name"] for tool in tools] == list(tool_names), tools
|
||||
messages: Final = body["messages"]
|
||||
assert isinstance(messages, list) and object_value(messages[-1])["content"] == _QUESTION, messages
|
||||
return body
|
||||
|
||||
|
||||
def _spend_statuses(model: str) -> list[JsonValue]:
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows('SELECT status FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)),
|
||||
lambda found: len(found) >= 1,
|
||||
seconds=70,
|
||||
)
|
||||
return [row["status"] for row in rows]
|
||||
|
||||
|
||||
def _openai_client(gateway: Gateway) -> openai.OpenAI:
|
||||
return openai.OpenAI(base_url=str(gateway.client.base_url) + "/v1", api_key=gateway.key, max_retries=0)
|
||||
|
||||
|
||||
def _async_openai_client(gateway: Gateway) -> openai.AsyncOpenAI:
|
||||
return openai.AsyncOpenAI(base_url=str(gateway.client.base_url) + "/v1", api_key=gateway.key, max_retries=0)
|
||||
|
||||
|
||||
def _raw_events(gateway: Gateway, body: dict[str, JsonValue]) -> list[dict[str, JsonValue]]:
|
||||
with gateway.client.stream(
|
||||
"POST",
|
||||
"/v1/responses",
|
||||
json={**body, "cache": _NO_CACHE},
|
||||
headers={"Authorization": f"Bearer {gateway.key}"},
|
||||
) as response:
|
||||
assert response.status_code == 200, response.read()
|
||||
return [
|
||||
_JSON_OBJECT.validate_json(line.removeprefix("data: "))
|
||||
for line in response.iter_lines()
|
||||
if line.startswith("data: {")
|
||||
]
|
||||
|
||||
|
||||
def _raw_added_items(events: Sequence[dict[str, JsonValue]]) -> Iterator[tuple[JsonValue, JsonValue]]:
|
||||
for event in events:
|
||||
if event["type"] == "response.output_item.added":
|
||||
yield object_value(event["item"])["type"], event["output_index"]
|
||||
|
||||
|
||||
def _raw_completed_output_types(events: Sequence[dict[str, JsonValue]]) -> list[JsonValue]:
|
||||
completed: Final = [event for event in events if event["type"] == "response.completed"]
|
||||
assert len(completed) == 1, [event["type"] for event in events]
|
||||
output: Final = object_value(completed[0]["response"])["output"]
|
||||
assert isinstance(output, list)
|
||||
return [object_value(item)["type"] for item in output]
|
||||
|
||||
|
||||
def test_openai_sdk_responses_stream_with_a_tool_call_only_reply_announces_no_message_item(gateway: Gateway) -> None:
|
||||
identity: Final = f"chatcmpl-bridge-{uuid.uuid4().hex}"
|
||||
reply: Final = _sse_reply(identity, (_tool_call_delta(),), "tool_calls")
|
||||
with _vllm_server(lambda _: reply) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _bridged_vllm_model(scenario, wire)
|
||||
with _openai_client(gateway).responses.stream(
|
||||
model=model,
|
||||
input=_QUESTION,
|
||||
tools=[_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool
|
||||
store=False,
|
||||
extra_body={"cache": _NO_CACHE},
|
||||
) as stream:
|
||||
events: Final = list(stream)
|
||||
final: Final = stream.get_final_response()
|
||||
added: Final = [
|
||||
(event.item.type, event.output_index) for event in events if event.type == "response.output_item.added"
|
||||
]
|
||||
assert added == [("function_call", 0)], [event.type for event in events]
|
||||
assert [item.type for item in final.output] == ["function_call"], final.output
|
||||
call: Final = final.output[0]
|
||||
assert call.type == "function_call" and call.name == "get_weather" and call.call_id == _CALL_ID
|
||||
assert json.loads(call.arguments) == _ARGUMENTS
|
||||
assert final.output_text == ""
|
||||
created: Final = [event for event in events if event.type == "response.created"]
|
||||
assert len(created) == 1 and created[0].response.id == final.id, [event.type for event in events]
|
||||
_only_streamed_chat(wire, "get_weather")
|
||||
assert _spend_statuses(model) == ["success"]
|
||||
|
||||
|
||||
async def test_async_openai_sdk_responses_stream_text_reply_has_one_message_item_at_index_zero(
|
||||
gateway: Gateway,
|
||||
) -> None:
|
||||
identity: Final = f"chatcmpl-bridge-{uuid.uuid4().hex}"
|
||||
reply: Final = _sse_reply(identity, ({"content": _ANSWER[:17]}, {"content": _ANSWER[17:]}), "stop")
|
||||
with _vllm_server(lambda _: reply) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _bridged_vllm_model(scenario, wire)
|
||||
stream: Final = await _async_openai_client(gateway).responses.create(
|
||||
model=model, input=_QUESTION, store=False, stream=True, extra_body={"cache": _NO_CACHE}
|
||||
)
|
||||
events: Final = [event async for event in stream]
|
||||
added: Final = [
|
||||
(event.item.type, event.output_index) for event in events if event.type == "response.output_item.added"
|
||||
]
|
||||
assert added == [("message", 0)], [event.type for event in events]
|
||||
done: Final = [(event.item.type, event.output_index) for event in events if event.type == "response.output_item.done"]
|
||||
assert done == [("message", 0)], [event.type for event in events]
|
||||
assert "".join(event.delta for event in events if event.type == "response.output_text.delta") == _ANSWER
|
||||
completed: Final = [event for event in events if event.type == "response.completed"]
|
||||
assert len(completed) == 1 and completed[0].response.output_text == _ANSWER
|
||||
assert [item.type for item in completed[0].response.output] == ["message"], completed[0].response.output
|
||||
created: Final = [event for event in events if event.type == "response.created"]
|
||||
assert len(created) == 1 and created[0].response.id == completed[0].response.id
|
||||
_only_streamed_chat(wire)
|
||||
assert _spend_statuses(model) == ["success"]
|
||||
|
||||
|
||||
async def test_async_openai_sdk_responses_stream_reasoning_then_text_gets_contiguous_output_indexes(
|
||||
gateway: Gateway,
|
||||
) -> None:
|
||||
identity: Final = f"chatcmpl-bridge-{uuid.uuid4().hex}"
|
||||
reply: Final = _sse_reply(identity, ({"reasoning_content": _REASONING}, {"content": _ANSWER}), "stop")
|
||||
with _vllm_server(lambda _: reply) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _bridged_vllm_model(scenario, wire)
|
||||
stream: Final = await _async_openai_client(gateway).responses.create(
|
||||
model=model, input=_QUESTION, store=False, stream=True, extra_body={"cache": _NO_CACHE}
|
||||
)
|
||||
events: Final = [event async for event in stream]
|
||||
added: Final = [
|
||||
(event.item.type, event.output_index) for event in events if event.type == "response.output_item.added"
|
||||
]
|
||||
assert added == [("reasoning", 0), ("message", 1)], [event.type for event in events]
|
||||
assert "".join(event.delta for event in events if event.type == "response.output_text.delta") == _ANSWER
|
||||
completed: Final = [event for event in events if event.type == "response.completed"]
|
||||
assert len(completed) == 1 and completed[0].response.output_text == _ANSWER
|
||||
assert [item.type for item in completed[0].response.output] == ["reasoning", "message"], (
|
||||
completed[0].response.output
|
||||
)
|
||||
_only_streamed_chat(wire)
|
||||
assert _spend_statuses(model) == ["success"]
|
||||
|
||||
|
||||
def test_raw_responses_stream_text_then_tool_call_keeps_the_message_first(gateway: Gateway) -> None:
|
||||
identity: Final = f"chatcmpl-bridge-{uuid.uuid4().hex}"
|
||||
reply: Final = _sse_reply(identity, ({"content": _PREFACE}, _tool_call_delta()), "tool_calls")
|
||||
with _vllm_server(lambda _: reply) as wire, gateway.scenario() as scenario:
|
||||
model: Final = _bridged_vllm_model(scenario, wire)
|
||||
events: Final = _raw_events(
|
||||
gateway, {"model": model, "input": _QUESTION, "tools": [_TOOL], "store": False, "stream": True}
|
||||
)
|
||||
assert list(_raw_added_items(events)) == [("message", 0), ("function_call", 1)], [e["type"] for e in events]
|
||||
assert _raw_completed_output_types(events) == ["message", "function_call"]
|
||||
deltas: Final = [event["delta"] for event in events if event["type"] == "response.output_text.delta"]
|
||||
assert "".join(str(delta) for delta in deltas) == _PREFACE, deltas
|
||||
done_calls: Final = [
|
||||
object_value(event["item"])
|
||||
for event in events
|
||||
if event["type"] == "response.output_item.done" and object_value(event["item"])["type"] == "function_call"
|
||||
]
|
||||
assert [(item["name"], item["call_id"]) for item in done_calls] == [("get_weather", _CALL_ID)], done_calls
|
||||
_only_streamed_chat(wire, "get_weather")
|
||||
assert _spend_statuses(model) == ["success"]
|
||||
|
|
@ -207,6 +207,84 @@ def test_ollama_pt_consecutive_user_messages():
|
|||
assert result["prompt"] == expected_prompt
|
||||
|
||||
|
||||
def _ollama_tool_turn(*results: object) -> list[dict]:
|
||||
call: Final = {"id": "call_1", "type": "function", "function": {"name": "get_weather", "arguments": '{"city": "Paris"}'}}
|
||||
return [
|
||||
{"role": "user", "content": "Weather in Paris?"},
|
||||
{"role": "assistant", "content": None, "tool_calls": [call]},
|
||||
*({"role": "tool", "tool_call_id": "call_1", "content": result} for result in results),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("results", "forwarded"),
|
||||
[
|
||||
pytest.param(("Paris: 22 degrees", "Sky: clear"), "Paris: 22 degrees\nSky: clear", id="two-tool-messages"),
|
||||
pytest.param(
|
||||
([{"type": "text", "text": "Paris: 22 degrees"}, {"type": "text", "text": "clear skies"}],),
|
||||
"Paris: 22 degrees\nclear skies",
|
||||
id="text-parts-of-one-tool-message",
|
||||
),
|
||||
pytest.param(
|
||||
([{"type": "text", "text": "Paris: 22 degrees"}], "Sky: clear"),
|
||||
"Paris: 22 degrees\nSky: clear",
|
||||
id="text-part-then-string",
|
||||
),
|
||||
pytest.param(
|
||||
([{"type": "text", "text": ""}, {"type": "text", "text": "clear skies"}],),
|
||||
"clear skies",
|
||||
id="empty-text-part-adds-no-blank-line",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_ollama_pt_separates_merged_tool_results_with_a_newline(results: tuple[object, ...], forwarded: str):
|
||||
result: Final = ollama_pt(model="llama2", messages=_ollama_tool_turn(*results))
|
||||
|
||||
assert isinstance(result, dict)
|
||||
assert result["prompt"].endswith(f"### User:\n{forwarded}\n\n"), result["prompt"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("content", [22, 22.5, True, {"temperature": 22}], ids=type)
|
||||
def test_ollama_pt_rejects_non_text_tool_content_as_a_bad_request(content: object):
|
||||
with pytest.raises(litellm.BadRequestError) as excinfo:
|
||||
ollama_pt(model="llama2", messages=_ollama_tool_turn(content))
|
||||
|
||||
assert excinfo.value.status_code == 400
|
||||
assert "content" in excinfo.value.message
|
||||
assert "tool message at index 2" in excinfo.value.message
|
||||
assert type(content).__name__ in excinfo.value.message
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("part", "expected_detail"),
|
||||
(
|
||||
({"type": "image_url", "image_url": None}, "NoneType image_url"),
|
||||
({"type": "text", "text": 22}, "int text part"),
|
||||
({"type": "text"}, "text part with no text"),
|
||||
({"type": "image_url"}, "image_url part with no image_url"),
|
||||
({"type": "image_url", "image_url": {"detail": "high"}}, "image_url object without a url string"),
|
||||
("hello", "str content part"),
|
||||
),
|
||||
ids=(
|
||||
"none-image-url",
|
||||
"int-text",
|
||||
"text-without-text",
|
||||
"image-url-without-image-url",
|
||||
"image-url-object-without-url",
|
||||
"str-part",
|
||||
),
|
||||
)
|
||||
def test_ollama_pt_rejects_a_malformed_content_part_as_a_bad_request(part: object, expected_detail: str):
|
||||
messages: Final = [{"role": "user", "content": [part]}]
|
||||
|
||||
with pytest.raises(litellm.BadRequestError) as excinfo:
|
||||
ollama_pt(model="llava", messages=messages)
|
||||
|
||||
assert excinfo.value.status_code == 400
|
||||
assert "user message at index 0" in excinfo.value.message
|
||||
assert expected_detail in excinfo.value.message
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_anthropic_bedrock_thinking_blocks_with_none_content():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -596,6 +596,28 @@ class TestOllamaConfig:
|
|||
|
||||
|
||||
class TestOllamaTextCompletionResponseIterator:
|
||||
def test_every_chunk_of_one_stream_carries_the_same_response_id(self):
|
||||
iterator: Final = OllamaTextCompletionResponseIterator(
|
||||
streaming_response=iter([]), sync_stream=True, json_mode=False
|
||||
)
|
||||
ollama_chunks: Final = (
|
||||
{"model": "qwen3:0.6b", "created_at": "2026-10-07T00:00:00Z", "response": "", "done": False},
|
||||
{"model": "qwen3:0.6b", "created_at": "2026-10-07T00:00:00Z", "response": "", "thinking": "Hm", "done": False},
|
||||
{"model": "qwen3:0.6b", "created_at": "2026-10-07T00:00:00Z", "response": "Hel", "done": False},
|
||||
{"model": "qwen3:0.6b", "created_at": "2026-10-07T00:00:00Z", "response": "lo", "done": False},
|
||||
)
|
||||
|
||||
results: Final = tuple(iterator.chunk_parser(chunk) for chunk in ollama_chunks)
|
||||
|
||||
ids: Final = {result.id for result in results if isinstance(result, ModelResponseStream)}
|
||||
assert len(results) == len(ollama_chunks) and len(ids) == 1, ids
|
||||
assert next(iter(ids)).startswith("chatcmpl-")
|
||||
other: Final = OllamaTextCompletionResponseIterator(
|
||||
streaming_response=iter([]), sync_stream=True, json_mode=False
|
||||
)
|
||||
other_result: Final = other.chunk_parser(ollama_chunks[2])
|
||||
assert isinstance(other_result, ModelResponseStream) and other_result.id not in ids
|
||||
|
||||
def test_chunk_parser_with_thinking_field(self):
|
||||
"""Test that chunks with 'thinking' field and empty 'response' are handled correctly."""
|
||||
iterator = OllamaTextCompletionResponseIterator(
|
||||
|
|
|
|||
|
|
@ -3954,7 +3954,9 @@ class TestEnsureOutputItemContentPartAdded:
|
|||
iterator._tool_item_id_by_call_id = {}
|
||||
iterator._tool_call_id_by_index = {}
|
||||
iterator._ambiguous_tool_call_indexes = set()
|
||||
iterator._next_tool_output_index = 1
|
||||
iterator._next_output_index = 0
|
||||
iterator._message_output_index = None
|
||||
iterator._reasoning_output_index = None
|
||||
iterator._final_tool_events_queued = False
|
||||
iterator._custom_tool_names = set()
|
||||
iterator.responses_api_request = {}
|
||||
|
|
@ -4456,7 +4458,7 @@ class TestEnsureOutputItemContentPartAdded:
|
|||
Choices(
|
||||
finish_reason=finish_reason,
|
||||
index=0,
|
||||
message=Message(content="", role="assistant"),
|
||||
message=Message(content="Partial answer", role="assistant"),
|
||||
)
|
||||
],
|
||||
usage=Usage(prompt_tokens=10, completion_tokens=1, total_tokens=11),
|
||||
|
|
|
|||
|
|
@ -132,14 +132,14 @@ def test_tool_call_delta_is_emitted_as_responses_events():
|
|||
evt1 = iterator._transform_chat_completion_chunk_to_response_api_chunk(chunk)
|
||||
assert evt1 is not None
|
||||
assert evt1.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED
|
||||
assert evt1.output_index == 1
|
||||
assert evt1.output_index == 0
|
||||
|
||||
# The arguments are now chunked, so we get the first delta chunk
|
||||
evt2 = iterator._transform_chat_completion_chunk_to_response_api_chunk(chunk)
|
||||
assert evt2 is not None
|
||||
assert evt2.type == ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA
|
||||
assert evt2.item_id == "fc_call_1"
|
||||
assert evt2.output_index == 1
|
||||
assert evt2.output_index == 0
|
||||
# The delta will be a chunk of the arguments, not the full arguments
|
||||
assert len(evt2.delta) <= 10 # Chunks are max 10 characters
|
||||
|
||||
|
|
@ -391,7 +391,7 @@ def test_tool_calls_present_only_in_final_response_are_emitted_before_completed(
|
|||
# First common_done_event_logic call should yield tool events, not response.completed.
|
||||
evt1 = iterator.common_done_event_logic(sync_mode=True)
|
||||
assert evt1.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED
|
||||
assert evt1.output_index == 1
|
||||
assert evt1.output_index == 0
|
||||
|
||||
# Now delta events are emitted (arguments split into chunks)
|
||||
# Collect all delta events
|
||||
|
|
@ -412,12 +412,12 @@ def test_tool_calls_present_only_in_final_response_are_emitted_before_completed(
|
|||
# The last event should be FUNCTION_CALL_ARGUMENTS_DONE
|
||||
assert evt.type == ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE
|
||||
assert evt.item_id == "fc_call_2"
|
||||
assert evt.output_index == 1
|
||||
assert evt.output_index == 0
|
||||
assert evt.arguments == '{"y":2}'
|
||||
|
||||
evt_final = iterator.common_done_event_logic(sync_mode=True)
|
||||
assert evt_final.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE
|
||||
assert evt_final.output_index == 1
|
||||
assert evt_final.output_index == 0
|
||||
|
||||
|
||||
def test_tool_call_arguments_are_chunked_to_match_openai_behavior():
|
||||
|
|
@ -472,7 +472,7 @@ def test_tool_call_arguments_are_chunked_to_match_openai_behavior():
|
|||
# First event should be OUTPUT_ITEM_ADDED
|
||||
assert evt is not None
|
||||
assert evt.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED
|
||||
assert evt.output_index == 1
|
||||
assert evt.output_index == 0
|
||||
assert hasattr(evt, "__dict__") and "sequence_number" in evt.__dict__
|
||||
|
||||
# Collect all remaining delta events from the pending queue by creating empty chunks
|
||||
|
|
@ -506,7 +506,7 @@ def test_tool_call_arguments_are_chunked_to_match_openai_behavior():
|
|||
for evt in delta_events:
|
||||
assert len(evt.delta) <= 10
|
||||
assert evt.item_id == "fc_call_test"
|
||||
assert evt.output_index == 1
|
||||
assert evt.output_index == 0
|
||||
assert hasattr(evt, "__dict__") and "sequence_number" in evt.__dict__
|
||||
|
||||
# Verify all deltas concatenated equal the original arguments
|
||||
|
|
@ -978,6 +978,23 @@ def _reasoning_chunk(reasoning: str, finish_reason: str | None = None) -> ModelR
|
|||
)
|
||||
|
||||
|
||||
def _annotation_only_chunk() -> ModelResponseStream:
|
||||
citation: Final = {"start_index": 0, "end_index": 2, "url": "https://example.com", "title": "Example"}
|
||||
return ModelResponseStream(
|
||||
id=CHAT_COMPLETION_ID,
|
||||
created=1748575031,
|
||||
model="claude-haiku-4-5",
|
||||
object="chat.completion.chunk",
|
||||
choices=[
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(role="assistant", annotations=[{"type": "url_citation", "url_citation": citation}]),
|
||||
finish_reason=None,
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def _signature_only_thinking_chunk(signature: str) -> ModelResponseStream:
|
||||
return ModelResponseStream(
|
||||
id=CHAT_COMPLETION_ID,
|
||||
|
|
@ -1034,6 +1051,68 @@ async def test_tool_only_stream_emits_no_message_item_events(sync_mode: bool):
|
|||
assert any(getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED for event in events)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.parametrize(
|
||||
"leading_chunk",
|
||||
[pytest.param(_chunk(""), id="empty-text-delta"), pytest.param(_reasoning_chunk(""), id="empty-reasoning-delta")],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_leading_delta_does_not_open_a_message_item_ahead_of_a_tool_call(
|
||||
sync_mode: bool, leading_chunk: ModelResponseStream
|
||||
):
|
||||
iterator: Final = _build_iterator([leading_chunk, _tool_call_chunk(), _chunk("", finish_reason="tool_calls")])
|
||||
|
||||
events: Final = await _collect_events(iterator, sync_mode)
|
||||
|
||||
item_events: Final = [
|
||||
(event.type, event.item.type)
|
||||
for event in events
|
||||
if getattr(event, "type", None)
|
||||
in (ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE)
|
||||
]
|
||||
assert item_events == [
|
||||
(ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, "function_call"),
|
||||
(ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, "function_call"),
|
||||
]
|
||||
completed: Final = [
|
||||
event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
|
||||
]
|
||||
assert [item.type for item in completed[0].response.output] == ["function_call"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_leading_delta_still_opens_the_message_item_for_the_first_text_delta(sync_mode: bool):
|
||||
iterator: Final = _build_iterator([_chunk(""), _chunk("Hi"), _chunk("", finish_reason="stop")])
|
||||
|
||||
events: Final = await _collect_events(iterator, sync_mode)
|
||||
|
||||
added_types: Final = [
|
||||
event.item.type
|
||||
for event in events
|
||||
if getattr(event, "type", None) == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED
|
||||
]
|
||||
assert added_types == ["message"]
|
||||
completed: Final = [
|
||||
event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
|
||||
]
|
||||
assert completed[0].response.output_text == "Hi"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_annotation_only_leading_delta_opens_the_message_item_before_its_annotation(sync_mode: bool):
|
||||
iterator: Final = _build_iterator([_annotation_only_chunk(), _chunk("Hi"), _chunk("", finish_reason="stop")])
|
||||
|
||||
events: Final = await _collect_events(iterator, sync_mode)
|
||||
|
||||
event_types: Final = [getattr(event, "type", None) for event in events]
|
||||
assert ResponsesAPIStreamEvents.OUTPUT_TEXT_ANNOTATION_ADDED in event_types
|
||||
assert event_types.index(ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED) < event_types.index(
|
||||
ResponsesAPIStreamEvents.OUTPUT_TEXT_ANNOTATION_ADDED
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_signature_only_thinking_streams_a_replayable_reasoning_item(sync_mode: bool):
|
||||
|
|
@ -1152,6 +1231,82 @@ async def test_tool_then_reasoning_then_text_gives_message_its_own_output_index(
|
|||
assert len(output_indexes_by_item_id) == len(set(output_indexes_by_item_id.values()))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_tool_only_stream_puts_the_function_call_at_output_index_zero(sync_mode: bool):
|
||||
iterator: Final = _build_iterator([_tool_call_chunk(), _chunk("", finish_reason="tool_calls")])
|
||||
|
||||
events: Final = await _collect_events(iterator, sync_mode)
|
||||
|
||||
indexed_events: Final = [event for event in events if hasattr(event, "output_index")]
|
||||
assert indexed_events
|
||||
assert {event.output_index for event in indexed_events} == {0}
|
||||
completed: Final = next(
|
||||
event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
|
||||
)
|
||||
assert [item.type for item in completed.response.output] == ["function_call"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_completed_output_follows_the_streamed_output_indexes(sync_mode: bool):
|
||||
iterator: Final = _build_iterator(
|
||||
[
|
||||
_reasoning_chunk("thinking"),
|
||||
_tool_call_chunk(),
|
||||
_chunk("Hello"),
|
||||
_chunk("!", finish_reason="stop"),
|
||||
]
|
||||
)
|
||||
|
||||
events: Final = await _collect_events(iterator, sync_mode)
|
||||
|
||||
added: Final = [
|
||||
event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED
|
||||
]
|
||||
assert [event.output_index for event in added] == list(range(len(added)))
|
||||
completed: Final = next(
|
||||
event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
|
||||
)
|
||||
assert [item.type for item in completed.response.output] == [event.item.type for event in added]
|
||||
assert [item.id for item in completed.response.output] == [event.item.id for event in added]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_reasoning_only_stream_lists_no_message_item_it_never_announced(sync_mode: bool):
|
||||
iterator: Final = _build_iterator([_reasoning_chunk("thinking"), _chunk("", finish_reason="stop")])
|
||||
|
||||
events: Final = await _collect_events(iterator, sync_mode)
|
||||
|
||||
added: Final = [
|
||||
event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED
|
||||
]
|
||||
assert [event.item.type for event in added] == ["reasoning"]
|
||||
completed: Final = next(
|
||||
event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
|
||||
)
|
||||
assert [item.type for item in completed.response.output] == ["reasoning"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_answer_stream_keeps_the_message_item_it_announced(sync_mode: bool):
|
||||
iterator: Final = _build_iterator([_chunk("", finish_reason="stop")])
|
||||
|
||||
events: Final = await _collect_events(iterator, sync_mode)
|
||||
|
||||
added: Final = [
|
||||
event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED
|
||||
]
|
||||
assert [event.item.type for event in added] == ["message"]
|
||||
completed: Final = next(
|
||||
event for event in events if getattr(event, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED
|
||||
)
|
||||
assert [item.type for item in completed.response.output] == ["message"]
|
||||
assert [item.id for item in completed.response.output] == [added[0].item.id]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sync_mode", [True, False])
|
||||
@pytest.mark.asyncio
|
||||
async def test_plain_text_stream_announces_exactly_one_message_item(sync_mode: bool):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue