From 141dba23721c85996119a23267003ca4e1872feb Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 9 Oct 2026 16:55:15 -0700 Subject: [PATCH] fix(proxy): keep the mapped status on assistants and threads route errors (#44994) * fix(proxy): keep the mapped status on assistants and threads route errors * test(proxy): type and annotate the assistants and threads route error tests * fix(proxy): redact internal details from assistants and threads route error messages * test(proxy): integration cells for assistants route error mapping * test(proxy): drop malformed JSON cases that never reach the assistants route handler --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/proxy/common_request_processing.py | 21 + litellm/proxy/proxy_server.py | 130 +---- .../_assistants_route_errors_support.py | 440 +++++++++++++++++ .../test_assistants_route_errors_chaos.py | 91 ++++ .../test_assistants_route_errors_wire.py | 445 ++++++++++++++++++ .../proxy_server/test_routes_assistants.py | 38 ++ .../proxy/proxy_server/test_routes_threads.py | 41 ++ .../proxy/test_common_request_processing.py | 49 ++ 8 files changed, 1134 insertions(+), 121 deletions(-) create mode 100644 tests/integration/providers/_assistants_route_errors_support.py create mode 100644 tests/integration/providers/test_assistants_route_errors_chaos.py create mode 100644 tests/integration/providers/test_assistants_route_errors_wire.py diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 317203cc2d5..4f18c773a24 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -636,6 +636,27 @@ def proxy_exception_from_http_exception(exc: HTTPException, headers: dict[str, s ) +def proxy_exception_from_route_error(exc: Exception) -> ProxyException: + if isinstance(exc, ProxyException): + return exc + if isinstance(exc, HTTPException): + return proxy_exception_from_http_exception(exc, {}) + error_status: Final = error_status_code(exc, status.HTTP_500_INTERNAL_SERVER_ERROR) + message: Final = attribute_of(exc, "message", str(exc)) + carried_code: Final = attribute_of(exc, "code") + provider_fields: Final = attribute_of(exc, "provider_specific_fields") + return ProxyException( + message=redact_internal_details_from_client_message( + strip_bug_report_notice(message) if isinstance(message, str) else str(exc) + ), + type=openai_error_type(exc, error_status), + param=openai_error_param(exc), + code=error_status, + openai_code=carried_code if isinstance(carried_code, str) else None, + provider_specific_fields=provider_fields if isinstance(provider_fields, dict) else None, + ) + + def _collect_response_file_search_vector_store_ids(data: Mapping[str, object]) -> set[str]: vector_store_ids: Final[set[str]] = set() tools: Final = data.get("tools") diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index e7226f62757..53f49f78c9f 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -408,6 +408,7 @@ from litellm.proxy.common_request_processing import ( # noqa: F401, RUF100 # l is_azure_model_router_request, log_llm_api_exception, open_sse_before_first_byte, + proxy_exception_from_route_error, request_litellm_call_id, resolve_litellm_call_id, should_return_raw_model_name, @@ -13429,22 +13430,7 @@ async def get_assistants( ) verbose_proxy_logger.error("litellm.proxy.proxy_server.get_assistants(): Exception occured - %s", e) verbose_proxy_logger.debug(traceback.format_exc()) - if isinstance(e, HTTPException): - raise ProxyException( - message=getattr(e, "message", str(e.detail)), - type=getattr(e, "type", "None"), - param=getattr(e, "param", "None"), - code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), - ) - else: - error_msg: Final = f"{e}" - raise ProxyException( - message=getattr(e, "message", error_msg), - type=getattr(e, "type", "None"), - param=getattr(e, "param", "None"), - openai_code=getattr(e, "code", None), - code=getattr(e, "status_code", 500), - ) + raise proxy_exception_from_route_error(e) @router.post( @@ -13520,21 +13506,7 @@ async def create_assistant( ) verbose_proxy_logger.error("litellm.proxy.proxy_server.create_assistant(): Exception occured - %s", e) verbose_proxy_logger.debug(traceback.format_exc()) - if isinstance(e, HTTPException): - raise ProxyException( - message=getattr(e, "message", str(e.detail)), - type=getattr(e, "type", "None"), - param=getattr(e, "param", "None"), - code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), - ) - else: - error_msg: Final = f"{e}" - raise ProxyException( - message=getattr(e, "message", error_msg), - type=getattr(e, "type", "None"), - param=getattr(e, "param", "None"), - code=getattr(e, "code", getattr(e, "status_code", 500)), - ) + raise proxy_exception_from_route_error(e) @router.delete( @@ -13609,21 +13581,7 @@ async def delete_assistant( ) verbose_proxy_logger.error("litellm.proxy.proxy_server.delete_assistant(): Exception occured - %s", e) verbose_proxy_logger.debug(traceback.format_exc()) - if isinstance(e, HTTPException): - raise ProxyException( - message=getattr(e, "message", str(e.detail)), - type=getattr(e, "type", "None"), - param=getattr(e, "param", "None"), - code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), - ) - else: - error_msg: Final = f"{e}" - raise ProxyException( - message=getattr(e, "message", error_msg), - type=getattr(e, "type", "None"), - param=getattr(e, "param", "None"), - code=getattr(e, "code", getattr(e, "status_code", 500)), - ) + raise proxy_exception_from_route_error(e) @router.post( @@ -13698,21 +13656,7 @@ async def create_threads( ) verbose_proxy_logger.error("litellm.proxy.proxy_server.create_threads(): Exception occured - %s", e) verbose_proxy_logger.debug(traceback.format_exc()) - if isinstance(e, HTTPException): - raise ProxyException( - message=getattr(e, "message", str(e.detail)), - type=getattr(e, "type", "None"), - param=getattr(e, "param", "None"), - code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), - ) - else: - error_msg: Final = f"{e}" - raise ProxyException( - message=getattr(e, "message", error_msg), - type=getattr(e, "type", "None"), - param=getattr(e, "param", "None"), - code=getattr(e, "code", getattr(e, "status_code", 500)), - ) + raise proxy_exception_from_route_error(e) @router.get( @@ -13785,21 +13729,7 @@ async def get_thread( ) verbose_proxy_logger.error("litellm.proxy.proxy_server.get_thread(): Exception occured - %s", e) verbose_proxy_logger.debug(traceback.format_exc()) - if isinstance(e, HTTPException): - raise ProxyException( - message=getattr(e, "message", str(e.detail)), - type=getattr(e, "type", "None"), - param=getattr(e, "param", "None"), - code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), - ) - else: - error_msg: Final = f"{e}" - raise ProxyException( - message=getattr(e, "message", error_msg), - type=getattr(e, "type", "None"), - param=getattr(e, "param", "None"), - code=getattr(e, "code", getattr(e, "status_code", 500)), - ) + raise proxy_exception_from_route_error(e) @router.post( @@ -13876,21 +13806,7 @@ async def add_messages( ) verbose_proxy_logger.error("litellm.proxy.proxy_server.add_messages(): Exception occured - %s", e) verbose_proxy_logger.debug(traceback.format_exc()) - if isinstance(e, HTTPException): - raise ProxyException( - message=getattr(e, "message", str(e.detail)), - type=getattr(e, "type", "None"), - param=getattr(e, "param", "None"), - code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), - ) - else: - error_msg: Final = f"{e}" - raise ProxyException( - message=getattr(e, "message", error_msg), - type=getattr(e, "type", "None"), - param=getattr(e, "param", "None"), - code=getattr(e, "code", getattr(e, "status_code", 500)), - ) + raise proxy_exception_from_route_error(e) @router.get( @@ -13963,21 +13879,7 @@ async def get_messages( ) verbose_proxy_logger.error("litellm.proxy.proxy_server.get_messages(): Exception occured - %s", e) verbose_proxy_logger.debug(traceback.format_exc()) - if isinstance(e, HTTPException): - raise ProxyException( - message=getattr(e, "message", str(e.detail)), - type=getattr(e, "type", "None"), - param=getattr(e, "param", "None"), - code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), - ) - else: - error_msg: Final = f"{e}" - raise ProxyException( - message=getattr(e, "message", error_msg), - type=getattr(e, "type", "None"), - param=getattr(e, "param", "None"), - code=getattr(e, "code", getattr(e, "status_code", 500)), - ) + raise proxy_exception_from_route_error(e) @router.post( @@ -14085,21 +13987,7 @@ async def run_thread( ) verbose_proxy_logger.error("litellm.proxy.proxy_server.run_thread(): Exception occured - %s", e) verbose_proxy_logger.debug(traceback.format_exc()) - if isinstance(e, HTTPException): - raise ProxyException( - message=getattr(e, "message", str(e.detail)), - type=getattr(e, "type", "None"), - param=getattr(e, "param", "None"), - code=getattr(e, "status_code", status.HTTP_400_BAD_REQUEST), - ) - else: - error_msg: Final = f"{e}" - raise ProxyException( - message=getattr(e, "message", error_msg), - type=getattr(e, "type", "None"), - param=getattr(e, "param", "None"), - code=getattr(e, "code", getattr(e, "status_code", 500)), - ) + raise proxy_exception_from_route_error(e) #### DEV UTILS #### diff --git a/tests/integration/providers/_assistants_route_errors_support.py b/tests/integration/providers/_assistants_route_errors_support.py new file mode 100644 index 00000000000..0a622818c2c --- /dev/null +++ b/tests/integration/providers/_assistants_route_errors_support.py @@ -0,0 +1,440 @@ +from __future__ import annotations + +import asyncio +import json +import re +import socket +import threading +import uuid +from collections.abc import AsyncIterator, Callable, Generator, Mapping +from contextlib import asynccontextmanager, contextmanager +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final + +import httpx +import psutil +import yaml +from integration._support.client import JSON_OBJECT, Gateway, eventually, object_value +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue + +MARKER: Final = re.compile(r"m[0-9a-f]{32}") +TRACEBACK: Final = "Traceback (most recent call last)" +BUG_NOTICE: Final = "This looks like a bug in LiteLLM" +ASSISTANT_MODEL: Final = "scripted-assistant-model" +CONTROL_MODEL: Final = "assistants-route-errors-control" +SESSION_HEADER: Final = MappingProxyType({"x-litellm-session-id": "assistants-route-errors"}) +STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") + + +@dataclass(frozen=True, slots=True) +class Route: + name: str + method: str + path: str + reads_body: bool + carries_marker: bool + + +ROUTES: Final = ( + Route("get_assistants", "GET", "/v1/assistants", reads_body=False, carries_marker=False), + Route("create_assistant", "POST", "/v1/assistants", reads_body=True, carries_marker=True), + Route("delete_assistant", "DELETE", "/v1/assistants/asst_{marker}", reads_body=False, carries_marker=True), + Route("create_thread", "POST", "/v1/threads", reads_body=False, carries_marker=False), + Route("get_thread", "GET", "/v1/threads/thread_{marker}", reads_body=False, carries_marker=True), + Route("add_message", "POST", "/v1/threads/thread_{marker}/messages", reads_body=True, carries_marker=True), + Route("get_messages", "GET", "/v1/threads/thread_{marker}/messages", reads_body=False, carries_marker=True), + Route("run_thread", "POST", "/v1/threads/thread_{marker}/runs", reads_body=True, carries_marker=True), +) +ROUTE_IDS: Final = tuple(route.name for route in ROUTES) +MARKED_ROUTES: Final = tuple(route for route in ROUTES if route.carries_marker) +BODY_ROUTES: Final = tuple(route for route in ROUTES if route.reads_body) +ROUTE_BY_NAME: Final = MappingProxyType({route.name: route for route in ROUTES}) + + +@dataclass(frozen=True, slots=True) +class Mapped: + status: int + type: str + param: str | None + + +STATUS_TABLE: Final = ( + Mapped(400, "scripted_type", "scripted_param"), + Mapped(401, "authentication_error", None), + Mapped(403, "permission_error", None), + Mapped(404, "invalid_request_error", None), + Mapped(408, "invalid_request_error", None), + Mapped(409, "invalid_request_error", None), + Mapped(422, "scripted_type", "scripted_param"), + Mapped(429, "throttling_error", "scripted_param"), + Mapped(500, "internal_server_error", "scripted_param"), + Mapped(502, "internal_server_error", None), + Mapped(503, "internal_server_error", None), + Mapped(504, "internal_server_error", None), +) +MAPPED_BY_STATUS: Final = MappingProxyType({mapped.status: mapped for mapped in STATUS_TABLE}) +UNMAPPED_500: Final = Mapped(500, "internal_server_error", None) + + +def new_marker() -> str: + return f"m{uuid.uuid4().hex}" + + +def free_port() -> int: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + return reserve.getsockname()[1] + + +def proxy_path(route: Route, marker: str) -> str: + return route.path.format(marker=marker) + + +def request_body(route: Route, marker: str) -> dict[str, JsonValue] | None: + match route.name: + case "create_assistant": + return {"model": ASSISTANT_MODEL, "name": marker} + case "create_thread": + return {"messages": [{"role": "user", "content": marker}]} + case "add_message": + return {"role": "user", "content": marker} + case "run_thread": + return {"assistant_id": f"asst_{marker}"} + case _: + return None + + +def marker_of(request: Request) -> str | None: + found: Final = MARKER.search(request.target) or MARKER.search(request.body.decode(errors="replace")) + return found.group(0) if found else None + + +def route_of(request: Request, marker: str) -> Route: + target: Final = (request.method, request.target.split("?", 1)[0]) + matches: Final = tuple(route for route in ROUTES if (route.method, proxy_path(route, marker)) == target) + assert len(matches) == 1, target + return matches[0] + + +def error_reply(status: int, message: str) -> Reply: + error: Final = {"message": message, "type": "scripted_type", "param": "scripted_param", "code": "scripted_code"} + return Reply(status=status, body=json.dumps({"error": error}).encode()) + + +def scripted_message(marker: str, status: int) -> str: + return f"scripted {marker} status {status}" + + +def _assistant(marker: str) -> dict[str, JsonValue]: + return { + "id": f"asst_{marker}", + "object": "assistant", + "created_at": 1, + "name": marker, + "description": None, + "model": ASSISTANT_MODEL, + "instructions": None, + "tools": [], + "metadata": {}, + } + + +def _thread(marker: str) -> dict[str, JsonValue]: + return {"id": f"thread_{marker}", "object": "thread", "created_at": 1, "metadata": {}, "tool_resources": None} + + +def _message(marker: str) -> dict[str, JsonValue]: + return { + "id": f"msg_{marker}", + "object": "thread.message", + "created_at": 1, + "thread_id": f"thread_{marker}", + "role": "user", + "content": [{"type": "text", "text": {"value": marker, "annotations": []}}], + "assistant_id": None, + "run_id": None, + "attachments": [], + "metadata": {}, + "status": "completed", + } + + +def _run(marker: str, status: str) -> dict[str, JsonValue]: + return { + "id": f"run_{marker}", + "object": "thread.run", + "created_at": 1, + "thread_id": f"thread_{marker}", + "assistant_id": f"asst_{marker}", + "status": status, + "model": ASSISTANT_MODEL, + "instructions": "", + "tools": [], + "metadata": {}, + "parallel_tool_calls": True, + } + + +def _page(item: dict[str, JsonValue]) -> dict[str, JsonValue]: + return {"object": "list", "data": [item], "first_id": item["id"], "last_id": item["id"], "has_more": False} + + +def success_body(route: Route, marker: str) -> dict[str, JsonValue]: + match route.name: + case "get_assistants": + return _page(_assistant(marker)) + case "create_assistant": + return _assistant(marker) + case "delete_assistant": + return {"id": f"asst_{marker}", "object": "assistant.deleted", "deleted": True} + case "create_thread" | "get_thread": + return _thread(marker) + case "add_message": + return _message(marker) + case "get_messages": + return _page(_message(marker)) + case _: + return _run(marker, "queued") + + +def run_poll_path(marker: str) -> str: + return f"/v1/threads/thread_{marker}/runs/run_{marker}" + + +def success_trail(route: Route, marker: str) -> tuple[tuple[str, str], ...]: + first: Final = (route.method, proxy_path(route, marker)) + return (first, ("GET", run_poll_path(marker))) if route.name == "run_thread" else (first,) + + +def chat_completion_body(marker: str) -> dict[str, JsonValue]: + return { + "id": f"chatcmpl-{marker}", + "object": "chat.completion", + "created": 1, + "model": "scripted-chat-model", + "choices": [{"index": 0, "message": {"role": "assistant", "content": marker}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + } + + +def success_peer(marker: str) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + if request.target == "/v1/chat/completions": + return Reply(body=json.dumps(chat_completion_body(marker)).encode()) + if request.method == "GET" and request.target == run_poll_path(marker): + return Reply(body=json.dumps(_run(marker, "completed")).encode()) + return Reply(body=json.dumps(success_body(route_of(request, marker), marker)).encode()) + + return respond + + +def error_object(response: httpx.Response) -> dict[str, JsonValue]: + body: Final = JSON_OBJECT.validate_json(response.content) + assert set(body) == {"error"}, response.text + return object_value(body["error"]) + + +def assert_clean_message(message: JsonValue) -> str: + assert isinstance(message, str), message + assert TRACEBACK not in message, message + assert BUG_NOTICE not in message, message + return message + + +def assert_openai_error(response: httpx.Response, mapped: Mapped) -> str: + assert response.status_code == mapped.status, response.text + error: Final = error_object(response) + assert set(error) == {"message", "type", "param", "code"}, response.text + assert (error["type"], error["param"], error["code"]) == (mapped.type, mapped.param, str(mapped.status)), ( + response.text + ) + return assert_clean_message(error["message"]) + + +def assert_mapped_upstream_error(response: httpx.Response, mapped: Mapped, marker: str) -> None: + message: Final = assert_openai_error(response, mapped) + assert marker in message, response.text + + +def is_model_listing(request: Request) -> bool: + return request.method == "GET" and request.target.split("?", 1)[0].endswith("/models") + + +def provider_requests(received: tuple[Request, ...]) -> tuple[Request, ...]: + return tuple(request for request in received if not is_model_listing(request)) + + +def upstream_trail(received: tuple[Request, ...]) -> tuple[tuple[str, str], ...]: + return tuple((request.method, request.target.split("?", 1)[0]) for request in provider_requests(received)) + + +@contextmanager +def assistants_wire(respond: Callable[[Request], Reply], port: int = 0) -> Generator[Wire, None, None]: + def answer(request: Request) -> Reply: + if is_model_listing(request): + return Reply(body=json.dumps({"object": "list", "data": []}).encode()) + return respond(request) + + with wire_server(answer, port=port) as wire: + yield wire + + +def assert_reached_upstream_once(received: tuple[Request, ...], route: Route, marker: str) -> None: + assert upstream_trail(received) == ((route.method, proxy_path(route, marker)),), received + + +def owned_config( + directory: Path, + model_list: tuple[Mapping[str, JsonValue], ...], + assistant_settings: Mapping[str, JsonValue], + general_settings: Mapping[str, JsonValue] = MappingProxyType({}), +) -> Path: + config: Final = JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + merged: Final = { + **config, + "model_list": [dict(model) for model in model_list], + "general_settings": {**object_value(config["general_settings"]), **general_settings}, + "router_settings": {**object_value(config["router_settings"]), "num_retries": 0}, + "assistant_settings": dict(assistant_settings), + } + path: Final = directory / f"assistants-route-errors-{uuid.uuid4().hex}.yaml" + path.write_text(yaml.safe_dump(merged)) + return path + + +def openai_assistants_config(directory: Path, upstream_port: int, timeout_seconds: int) -> Path: + api_base: Final = f"http://127.0.0.1:{upstream_port}/v1" + deployment: Final[dict[str, JsonValue]] = { + "api_base": api_base, + "api_key": "sk-scripted-assistants", + "max_retries": 0, + "timeout": timeout_seconds, + } + control: Final[dict[str, JsonValue]] = { + "model_name": CONTROL_MODEL, + "litellm_params": {"model": "openai/scripted-chat-model", **deployment}, + } + return owned_config(directory, (control,), {"custom_llm_provider": "openai", "litellm_params": deployment}) + + +def worker_pids(log: Path) -> tuple[int, ...]: + return tuple(int(pid) for pid in STARTED_WORKER.findall(log.read_text())) + + +def live_worker_pids(log: Path) -> tuple[int, ...]: + return tuple(pid for pid in worker_pids(log) if psutil.pid_exists(pid)) + + +def open_upstream_connections(pid: int, upstream_port: int) -> int: + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == upstream_port + ) + + +def held_upstream_connections(workers: tuple[int, ...], upstream_port: int, expected: int) -> Mapping[int, int]: + return eventually( + lambda: MappingProxyType({pid: open_upstream_connections(pid, upstream_port) for pid in workers}), + lambda held_by: sum(held_by.values()) == expected, + seconds=10, + ) + + +def held_error_peer( + status_by_marker: Mapping[str, int], arrived: SimpleQueue[str], release: threading.Event +) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + marker: Final = marker_of(request) + assert marker is not None, request.target + arrived.put(marker) + assert release.wait(timeout=60), "The burst was never released" + return error_reply(status_by_marker[marker], scripted_message(marker, status_by_marker[marker])) + + return respond + + +@dataclass(frozen=True, slots=True) +class Call: + route: Route + marker: str + + +@dataclass(frozen=True, slots=True) +class Answer: + call: Call + response: httpx.Response + + +async def _answer(client: httpx.AsyncClient, call: Call) -> Answer: + response: Final = await client.request( + call.route.method, proxy_path(call.route, call.marker), json=request_body(call.route, call.marker) + ) + return Answer(call, response) + + +def _async_client(gateway: Gateway) -> httpx.AsyncClient: + return httpx.AsyncClient( + base_url=str(gateway.client.base_url), + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=60, + trust_env=False, + ) + + +def _answers_of(results: tuple[Answer | BaseException, ...]) -> tuple[Answer, ...]: + for result in results: + assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result) + return tuple(result for result in results if isinstance(result, Answer)) + + +async def answer_all( + gateway: Gateway, calls: tuple[Call, ...], *, tolerate_transport_errors: bool = False +) -> tuple[Answer, ...]: + async with _async_client(gateway) as client: + results: Final = await asyncio.gather( + *(_answer(client, call) for call in calls), return_exceptions=tolerate_transport_errors + ) + return _answers_of(tuple(results)) + + +async def _send_one_by_one( + client: httpx.AsyncClient, + calls: tuple[Call, ...], + arrived: SimpleQueue[str], + enough: Callable[[], bool], +) -> tuple[asyncio.Task[Answer], ...]: + if not calls: + return () + sent: Final = asyncio.create_task(_answer(client, calls[0])) + assert await asyncio.to_thread(arrived.get, True, 30) == calls[0].marker + if enough(): + return (sent,) + return (sent, *await _send_one_by_one(client, calls[1:], arrived, enough)) + + +@dataclass(frozen=True, slots=True) +class Held: + calls: tuple[Call, ...] + pending: tuple[asyncio.Task[Answer], ...] + + async def answers(self) -> tuple[Answer, ...]: + return _answers_of(tuple(await asyncio.gather(*self.pending, return_exceptions=True))) + + +@asynccontextmanager +async def held_one_by_one( + gateway: Gateway, calls: tuple[Call, ...], arrived: SimpleQueue[str], enough: Callable[[], bool] +) -> AsyncIterator[Held]: + async with _async_client(gateway) as client: + pending: Final = await _send_one_by_one(client, calls, arrived, enough) + yield Held(calls[: len(pending)], pending) + + +def assert_answered_with_its_own_marker(answer: Answer, mapped: Mapped) -> None: + assert_mapped_upstream_error(answer.response, mapped, answer.call.marker) + assert set(MARKER.findall(answer.response.text)) == {answer.call.marker}, answer.response.text diff --git a/tests/integration/providers/test_assistants_route_errors_chaos.py b/tests/integration/providers/test_assistants_route_errors_chaos.py new file mode 100644 index 00000000000..5e863932050 --- /dev/null +++ b/tests/integration/providers/test_assistants_route_errors_chaos.py @@ -0,0 +1,91 @@ +import asyncio +import signal +import threading +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final +from urllib.parse import urlsplit + +import psutil +import pytest +from integration._support.client import Gateway, eventually +from integration._support.process import graceful_stop_seconds, owned_proxy_process +from integration.providers._assistants_route_errors_support import ( + MAPPED_BY_STATUS, + MARKED_ROUTES, + ROUTE_BY_NAME, + Call, + answer_all, + assert_answered_with_its_own_marker, + assistants_wire, + held_error_peer, + held_one_by_one, + held_upstream_connections, + live_worker_pids, + marker_of, + new_marker, + open_upstream_connections, + openai_assistants_config, + provider_requests, + worker_pids, +) + +_STATUSES: Final = (400, 404, 429, 500, 503) +_MIN_CALLS: Final = 20 +_MAX_CALLS: Final = 60 +_STARTUP_COMPLETE: Final = "Application startup complete." + + +def _startups(log: Path) -> tuple[int, int]: + return len(worker_pids(log)), log.read_text().count(_STARTUP_COMPLETE) + + +@pytest.mark.timeout(int(2 * graceful_stop_seconds() + 120)) +async def test_worker_sigkill_mid_burst_leaves_the_sibling_answering_each_mapped_status( + gateway: Gateway, tmp_path: Path +) -> None: + calls: Final = tuple(Call(MARKED_ROUTES[index % len(MARKED_ROUTES)], new_marker()) for index in range(_MAX_CALLS)) + follow_up: Final = Call(ROUTE_BY_NAME["get_thread"], new_marker()) + status_by_marker: Final = MappingProxyType( + { + **{call.marker: _STATUSES[index % len(_STATUSES)] for index, call in enumerate(calls)}, + follow_up.marker: 404, + } + ) + release: Final = threading.Event() + arrived: Final[SimpleQueue[str]] = SimpleQueue() + with assistants_wire(held_error_peer(status_by_marker, arrived, release)) as wire: + port: Final = urlsplit(wire.url).port + assert port is not None + config: Final = openai_assistants_config(tmp_path, port, 60) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + eventually(lambda: _startups(owned.log), lambda found: found == (2, 2), seconds=120) + workers: Final = live_worker_pids(owned.log) + assert len(workers) == 2, workers + + def both_workers_hold_enough() -> bool: + held: Final = tuple(open_upstream_connections(pid, port) for pid in workers) + return sum(held) >= _MIN_CALLS and min(held) > 0 + + try: + async with held_one_by_one(owned.gateway, calls, arrived, both_workers_hold_enough) as held: + held_by: Final = await asyncio.to_thread(held_upstream_connections, workers, port, len(held.calls)) + assert min(held_by.values()) > 0, held_by + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + release.set() + served: Final = await held.answers() + finally: + release.set() + (answered,) = await answer_all(owned.gateway, (follow_up,)) + await asyncio.to_thread(eventually, lambda: _startups(owned.log), lambda found: found == (3, 3), 180) + received: Final = provider_requests(wire.drain()) + assert len(served) == held_by[survivor_pid], (held_by, len(served)) + for answer in (*served, answered): + assert_answered_with_its_own_marker(answer, MAPPED_BY_STATUS[status_by_marker[answer.call.marker]]) + assert sorted(marker_of(request) or "" for request in received) == sorted( + (*(call.marker for call in held.calls), follow_up.marker) + ), received diff --git a/tests/integration/providers/test_assistants_route_errors_wire.py b/tests/integration/providers/test_assistants_route_errors_wire.py new file mode 100644 index 00000000000..ec6d58e6670 --- /dev/null +++ b/tests/integration/providers/test_assistants_route_errors_wire.py @@ -0,0 +1,445 @@ +import asyncio +import json +import threading +from collections.abc import Iterator, Mapping +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final + +import httpx +import openai +import pytest +from integration._support.client import Gateway, gateway_from_environment +from integration._support.process import graceful_stop_seconds, owned_proxy_process +from integration._support.wire import Reply, Request +from integration.providers._assistants_route_errors_support import ( + ASSISTANT_MODEL, + BODY_ROUTES, + CONTROL_MODEL, + MAPPED_BY_STATUS, + MARKED_ROUTES, + MARKER, + ROUTE_BY_NAME, + ROUTE_IDS, + ROUTES, + SESSION_HEADER, + STATUS_TABLE, + UNMAPPED_500, + Call, + Mapped, + Route, + answer_all, + assert_answered_with_its_own_marker, + assert_clean_message, + assert_mapped_upstream_error, + assert_openai_error, + assert_reached_upstream_once, + assistants_wire, + error_reply, + free_port, + held_error_peer, + held_one_by_one, + held_upstream_connections, + live_worker_pids, + new_marker, + openai_assistants_config, + owned_config, + provider_requests, + proxy_path, + request_body, + scripted_message, + success_peer, + success_trail, + upstream_trail, +) +from pydantic import JsonValue + +pytestmark = pytest.mark.timeout(int(2 * graceful_stop_seconds() + 120)) + +_DEPLOYMENT_TIMEOUT_SECONDS: Final = 15 +_AZURE_CREDENTIALS: Final = ( + "AZURE_API_KEY", + "AZURE_OPENAI_API_KEY", + "AZURE_AD_TOKEN", + "AZURE_OPENAI_AD_TOKEN", + "AZURE_CLIENT_ID", + "AZURE_CLIENT_SECRET", + "AZURE_TENANT_ID", +) +_PRIVATE_DETAILS: Final = ("10.20.30.40", "/etc/litellm/secrets/db.yaml", "sk-proj-" + "a1B2c3D4" * 5) +_URL_MODEL: Final = "http://169.254.169.254/latest" +_BURST_STATUSES: Final = (400, 401, 404, 429, 500, 503) +_SDK_ERRORS: Final = MappingProxyType({404: openai.NotFoundError, 429: openai.RateLimitError}) +_BODY_VARIANTS: Final = ("text_502", "empty_503", "hostile_400", "null_message_404", "garbage_200") + + +@dataclass(frozen=True, slots=True) +class _Scripted: + gateway: Gateway + port: int + log: Path + + +@pytest.fixture(scope="module") +def scripted(tmp_path_factory: pytest.TempPathFactory) -> Iterator[_Scripted]: + directory: Final = tmp_path_factory.mktemp("assistants-route-errors-scripted") + port: Final = free_port() + config: Final = openai_assistants_config(directory, port, _DEPLOYMENT_TIMEOUT_SECONDS) + with ( + gateway_from_environment() as shared, + owned_proxy_process(shared, directory, {}, config=config, workers=2) as owned, + ): + yield _Scripted(owned.gateway, port, owned.log) + + +@pytest.fixture(scope="module") +def azure_without_credentials(tmp_path_factory: pytest.TempPathFactory) -> Iterator[Gateway]: + directory: Final = tmp_path_factory.mktemp("assistants-route-errors-azure") + settings: Final[dict[str, JsonValue]] = { + "custom_llm_provider": "azure", + "litellm_params": {"api_base": "http://127.0.0.1:9/", "api_version": "2024-05-01-preview", "max_retries": 0}, + } + config: Final = owned_config(directory, (), settings, {"missing_session_id": "reject"}) + with ( + gateway_from_environment() as shared, + owned_proxy_process( + shared, directory, {}, config=config, remove_environment=_AZURE_CREDENTIALS, workers=2 + ) as owned, + ): + yield owned.gateway + + +def _send( + gateway: Gateway, route: Route, marker: str, headers: Mapping[str, str] = MappingProxyType({}) +) -> httpx.Response: + return gateway.request(route.method, proxy_path(route, marker), request_body(route, marker), headers=headers) + + +def _send_raw(gateway: Gateway, route: Route, marker: str, content: bytes) -> httpx.Response: + return gateway.client.request( + route.method, + proxy_path(route, marker), + content=content, + headers={"Authorization": f"Bearer {gateway.key}", "content-type": "application/json"}, + ) + + +def _scripted_error(status: int, marker: str) -> Reply: + return error_reply(status, scripted_message(marker, status)) + + +@pytest.mark.parametrize("status", [mapped.status for mapped in STATUS_TABLE]) +@pytest.mark.parametrize("route", ROUTES, ids=ROUTE_IDS) +def test_route_keeps_the_mapped_upstream_status(scripted: _Scripted, route: Route, status: int) -> None: + marker: Final = new_marker() + with assistants_wire(lambda _: _scripted_error(status, marker), port=scripted.port) as wire: + response: Final = _send(scripted.gateway, route, marker) + assert_mapped_upstream_error(response, MAPPED_BY_STATUS[status], marker) + assert_reached_upstream_once(wire.drain(), route, marker) + + +def _sdk_call(client: openai.OpenAI, route: Route, marker: str) -> object: + match route.name: + case "get_assistants": + return client.beta.assistants.list() + case "create_assistant": + return client.beta.assistants.create(model=ASSISTANT_MODEL, name=marker) + case "delete_assistant": + return client.beta.assistants.delete(f"asst_{marker}") + case "create_thread": + return client.beta.threads.create(messages=[{"role": "user", "content": marker}]) + case "get_thread": + return client.beta.threads.retrieve(f"thread_{marker}") + case "add_message": + return client.beta.threads.messages.create(f"thread_{marker}", role="user", content=marker) + case "get_messages": + return client.beta.threads.messages.list(f"thread_{marker}") + case _: + return client.beta.threads.runs.create(f"thread_{marker}", assistant_id=f"asst_{marker}") + + +async def _async_sdk_call(client: openai.AsyncOpenAI, route: Route, marker: str) -> object: + match route.name: + case "get_assistants": + return await client.beta.assistants.list() + case "create_assistant": + return await client.beta.assistants.create(model=ASSISTANT_MODEL, name=marker) + case "delete_assistant": + return await client.beta.assistants.delete(f"asst_{marker}") + case "create_thread": + return await client.beta.threads.create(messages=[{"role": "user", "content": marker}]) + case "get_thread": + return await client.beta.threads.retrieve(f"thread_{marker}") + case "add_message": + return await client.beta.threads.messages.create(f"thread_{marker}", role="user", content=marker) + case "get_messages": + return await client.beta.threads.messages.list(f"thread_{marker}") + case _: + return await client.beta.threads.runs.create(f"thread_{marker}", assistant_id=f"asst_{marker}") + + +def _sdk_base_url(gateway: Gateway) -> str: + return f"{str(gateway.client.base_url).rstrip('/')}/v1" + + +def _assert_sdk_error(error: openai.APIStatusError, mapped: Mapped, marker: str) -> None: + assert (error.status_code, error.code, error.type, error.param) == ( + mapped.status, + str(mapped.status), + mapped.type, + mapped.param, + ), error.message + assert marker in error.message, error.message + assert_clean_message(error.message) + + +@pytest.mark.parametrize("status", sorted(_SDK_ERRORS)) +@pytest.mark.parametrize("route", ROUTES, ids=ROUTE_IDS) +def test_sync_sdk_raises_the_mapped_error(scripted: _Scripted, route: Route, status: int) -> None: + marker: Final = new_marker() + with ( + assistants_wire(lambda _: _scripted_error(status, marker), port=scripted.port) as wire, + openai.OpenAI( + base_url=_sdk_base_url(scripted.gateway), api_key=scripted.gateway.key, max_retries=0, timeout=30 + ) as client, + ): + with pytest.raises(_SDK_ERRORS[status]) as raised: + _sdk_call(client, route, marker) + assert_reached_upstream_once(wire.drain(), route, marker) + _assert_sdk_error(raised.value, MAPPED_BY_STATUS[status], marker) + + +@pytest.mark.parametrize("status", sorted(_SDK_ERRORS)) +@pytest.mark.parametrize("route", ROUTES, ids=ROUTE_IDS) +async def test_async_sdk_raises_the_mapped_error(scripted: _Scripted, route: Route, status: int) -> None: + marker: Final = new_marker() + with assistants_wire(lambda _: _scripted_error(status, marker), port=scripted.port) as wire: + async with openai.AsyncOpenAI( + base_url=_sdk_base_url(scripted.gateway), api_key=scripted.gateway.key, max_retries=0, timeout=30 + ) as client: + with pytest.raises(_SDK_ERRORS[status]) as raised: + await _async_sdk_call(client, route, marker) + assert_reached_upstream_once(wire.drain(), route, marker) + _assert_sdk_error(raised.value, MAPPED_BY_STATUS[status], marker) + + +@pytest.mark.parametrize("route", ROUTES, ids=ROUTE_IDS) +def test_internal_details_in_an_upstream_message_are_redacted(scripted: _Scripted, route: Route) -> None: + marker: Final = new_marker() + upstream_message: Final = f"{scripted_message(marker, 500)} at {' '.join(_PRIVATE_DETAILS)}" + with assistants_wire(lambda _: error_reply(500, upstream_message), port=scripted.port) as wire: + response: Final = _send(scripted.gateway, route, marker) + assert_reached_upstream_once(wire.drain(), route, marker) + message: Final = assert_openai_error(response, MAPPED_BY_STATUS[500]) + assert marker in message and "REDACTED" in message, response.text + assert [detail for detail in _PRIVATE_DETAILS if detail in response.text] == [], response.text + + +@pytest.mark.parametrize("route", ROUTES, ids=ROUTE_IDS) +def test_azure_without_credentials_answers_a_clean_connection_error( + azure_without_credentials: Gateway, route: Route +) -> None: + response: Final = _send(azure_without_credentials, route, new_marker(), SESSION_HEADER) + message: Final = assert_openai_error(response, UNMAPPED_500) + assert "AzureException" in message, response.text + + +@pytest.mark.parametrize("route", ROUTES, ids=ROUTE_IDS) +def test_refused_upstream_answers_a_connection_error(scripted: _Scripted, route: Route) -> None: + response: Final = _send(scripted.gateway, route, new_marker()) + message: Final = assert_openai_error(response, UNMAPPED_500) + assert "litellm.APIConnectionError" in message, response.text + + +async def test_upstream_held_past_the_deployment_timeout_answers_408_on_every_route(scripted: _Scripted) -> None: + marker: Final = new_marker() + release: Final = threading.Event() + arrived: Final[SimpleQueue[str]] = SimpleQueue() + + def held(request: Request) -> Reply: + arrived.put(request.target) + assert release.wait(timeout=60), "The held requests were never released" + return _scripted_error(404, marker) + + with assistants_wire(held, port=scripted.port) as wire: + try: + answers: Final = await answer_all(scripted.gateway, tuple(Call(route, marker) for route in ROUTES)) + finally: + release.set() + received: Final = wire.drain() + assert len(answers) == len(ROUTES) + for answer in answers: + assert "litellm.Timeout" in assert_openai_error(answer.response, MAPPED_BY_STATUS[408]), answer.response.text + assert sorted(upstream_trail(received)) == sorted((route.method, proxy_path(route, marker)) for route in ROUTES) + + +_PLAIN_EXCEPTION_CASES: Final = (("run_thread", "assistant_id"), ("add_message", "role")) + + +@pytest.mark.parametrize( + ("route_name", "named"), _PLAIN_EXCEPTION_CASES, ids=("run_without_assistant", "message_without_role") +) +def test_plain_exception_answers_an_openai_shaped_500(scripted: _Scripted, route_name: str, named: str) -> None: + marker: Final = new_marker() + with assistants_wire(lambda _: _scripted_error(404, marker), port=scripted.port) as wire: + response: Final = _send_raw(scripted.gateway, ROUTE_BY_NAME[route_name], marker, b"{}") + assert provider_requests(wire.drain()) == () + assert named in assert_openai_error(response, UNMAPPED_500), response.text + + +@pytest.mark.parametrize("route", BODY_ROUTES, ids=[route.name for route in BODY_ROUTES]) +def test_url_valued_model_answers_like_chat_completions(scripted: _Scripted, route: Route) -> None: + marker: Final = new_marker() + body: Final[dict[str, JsonValue]] = {**(request_body(route, marker) or {}), "model": _URL_MODEL} + with assistants_wire(lambda _: _scripted_error(404, marker), port=scripted.port) as wire: + assistants: Final = scripted.gateway.request(route.method, proxy_path(route, marker), body) + chat: Final = scripted.gateway.request( + "POST", "/v1/chat/completions", {"model": _URL_MODEL, "messages": [{"role": "user", "content": marker}]} + ) + assert provider_requests(wire.drain()) == () + assert chat.status_code == 400, chat.text + assert (assistants.status_code, assistants.json()) == (chat.status_code, chat.json()), assistants.text + + +@pytest.mark.parametrize("route", ROUTES, ids=ROUTE_IDS) +def test_missing_session_id_rejection_keeps_its_400(azure_without_credentials: Gateway, route: Route) -> None: + response: Final = _send(azure_without_credentials, route, new_marker()) + message: Final = assert_openai_error(response, Mapped(400, "bad_request_error", "session_id")) + assert "session id" in message, response.text + + +@pytest.mark.parametrize("route", ROUTES, ids=ROUTE_IDS) +def test_missing_assistant_settings_answers_an_openai_shaped_500(gateway: Gateway, route: Route) -> None: + response: Final = _send(gateway, route, new_marker()) + message: Final = assert_openai_error(response, UNMAPPED_500) + assert "custom_llm_provider" in message, response.text + + +def _variant_reply(variant: str, marker: str) -> Reply: + match variant: + case "text_502": + return Reply( + status=502, content_type="text/plain", body=f"upstream proxy 10.0.0.7 failed {marker}".encode() + ) + case "empty_503": + return Reply(status=503, body=b"") + case "hostile_400": + hostile: Final = {"message": marker + "y" * 5000, "type": ["scripted"], "param": {"field": 1}, "code": 123} + return Reply(status=400, body=json.dumps({"error": hostile}).encode()) + case "null_message_404": + nulls: Final = {"message": None, "type": "scripted_type", "param": None, "code": None, "detail": marker} + return Reply(status=404, body=json.dumps({"error": nulls}).encode()) + case _: + return Reply(status=200, body=f"not json {marker}".encode()) + + +_VARIANT_MAPPED: Final = MappingProxyType( + { + "text_502": Mapped(502, "internal_server_error", None), + "empty_503": Mapped(503, "internal_server_error", None), + "hostile_400": Mapped(400, "invalid_request_error", None), + "null_message_404": Mapped(404, "invalid_request_error", None), + "garbage_200": UNMAPPED_500, + } +) +_VARIANTS_CARRYING_THE_MARKER: Final = frozenset(("text_502", "hostile_400", "null_message_404")) + + +@pytest.mark.parametrize("variant", _BODY_VARIANTS) +@pytest.mark.parametrize("route", ROUTES, ids=ROUTE_IDS) +def test_unusual_upstream_bodies_keep_status_and_string_fields(scripted: _Scripted, route: Route, variant: str) -> None: + marker: Final = new_marker() + with assistants_wire(lambda _: _variant_reply(variant, marker), port=scripted.port) as wire: + response: Final = _send(scripted.gateway, route, marker) + assert_reached_upstream_once(wire.drain(), route, marker) + message: Final = assert_openai_error(response, _VARIANT_MAPPED[variant]) + assert variant not in _VARIANTS_CARRYING_THE_MARKER or marker in message, response.text + assert "10.0.0.7" not in response.text, response.text + + +@pytest.mark.parametrize("route", ROUTES, ids=ROUTE_IDS) +def test_repeated_failing_request_answers_the_same_each_time(scripted: _Scripted, route: Route) -> None: + marker: Final = new_marker() + with assistants_wire(lambda _: _scripted_error(404, marker), port=scripted.port) as wire: + first: Final = _send(scripted.gateway, route, marker) + second: Final = _send(scripted.gateway, route, marker) + received: Final = wire.drain() + assert_mapped_upstream_error(first, MAPPED_BY_STATUS[404], marker) + assert first.json() == second.json(), (first.text, second.text) + assert upstream_trail(received) == ((route.method, proxy_path(route, marker)),) * 2, received + + +async def test_mixed_status_burst_answers_each_request_with_its_own_status_and_marker(scripted: _Scripted) -> None: + calls: Final = tuple(Call(MARKED_ROUTES[index % len(MARKED_ROUTES)], new_marker()) for index in range(48)) + status_by_marker: Final = MappingProxyType( + { + call.marker: _BURST_STATUSES[(index // len(MARKED_ROUTES)) % len(_BURST_STATUSES)] + for index, call in enumerate(calls) + } + ) + release: Final = threading.Event() + arrived: Final[SimpleQueue[str]] = SimpleQueue() + workers: Final = live_worker_pids(scripted.log) + assert len(workers) == 2, workers + with assistants_wire(held_error_peer(status_by_marker, arrived, release), port=scripted.port) as wire: + try: + async with held_one_by_one(scripted.gateway, calls, arrived, lambda: False) as held: + held_by: Final = await asyncio.to_thread(held_upstream_connections, workers, scripted.port, len(calls)) + release.set() + answers: Final = await held.answers() + finally: + release.set() + received: Final = wire.drain() + assert sum(held_by.values()) == len(calls) and min(held_by.values()) > 0, held_by + assert len(answers) == len(calls) + for answer in answers: + assert_answered_with_its_own_marker(answer, MAPPED_BY_STATUS[status_by_marker[answer.call.marker]]) + assert sorted(upstream_trail(received)) == sorted( + (call.route.method, proxy_path(call.route, call.marker)) for call in calls + ) + + +async def test_outage_then_recovery_answers_connection_errors_then_upstream_bodies(scripted: _Scripted) -> None: + marker: Final = new_marker() + refused: Final = await answer_all(scripted.gateway, tuple(Call(route, marker) for route in ROUTES)) + for answer in refused: + assert "litellm.APIConnectionError" in assert_openai_error(answer.response, UNMAPPED_500), answer.response.text + with assistants_wire(success_peer(marker), port=scripted.port) as wire: + recovered: Final = tuple(_send(scripted.gateway, route, marker) for route in ROUTES) + chat: Final = scripted.gateway.request( + "POST", "/v1/chat/completions", {"model": CONTROL_MODEL, "messages": [{"role": "user", "content": marker}]} + ) + received: Final = wire.drain() + for route, response in zip(ROUTES, recovered, strict=True): + assert response.status_code == 200, (route.name, response.text) + assert marker in response.text, (route.name, response.text) + assert chat.status_code == 200 and marker in chat.text, chat.text + assert upstream_trail(received) == ( + *(step for route in ROUTES for step in success_trail(route, marker)), + ("POST", "/v1/chat/completions"), + ), received + + +@pytest.mark.parametrize("route", ROUTES, ids=ROUTE_IDS) +def test_control_upstream_success_answers_200_with_the_upstream_body(scripted: _Scripted, route: Route) -> None: + marker: Final = new_marker() + with assistants_wire(success_peer(marker), port=scripted.port) as wire: + response: Final = _send(scripted.gateway, route, marker) + received: Final = wire.drain() + assert response.status_code == 200, response.text + assert marker in response.text, response.text + assert upstream_trail(received) == success_trail(route, marker), received + + +def test_control_streamed_run_keeps_the_upstream_404(scripted: _Scripted) -> None: + marker: Final = new_marker() + route: Final = ROUTE_BY_NAME["run_thread"] + with assistants_wire(lambda _: _scripted_error(404, marker), port=scripted.port) as wire: + response: Final = scripted.gateway.request( + "POST", proxy_path(route, marker), {"assistant_id": f"asst_{marker}", "stream": True} + ) + assert_reached_upstream_once(wire.drain(), route, marker) + assert response.status_code == 404, response.text + assert MARKER.search(response.text) is not None, response.text diff --git a/tests/unit/proxy/proxy_server/test_routes_assistants.py b/tests/unit/proxy/proxy_server/test_routes_assistants.py index fd1f672dce9..2995978aecb 100644 --- a/tests/unit/proxy/proxy_server/test_routes_assistants.py +++ b/tests/unit/proxy/proxy_server/test_routes_assistants.py @@ -11,10 +11,14 @@ Pins (PR2): from __future__ import annotations +from contextlib import AbstractContextManager +from typing import Callable, Final from unittest.mock import AsyncMock, MagicMock import pytest +from fastapi.testclient import TestClient +import litellm from litellm.proxy import proxy_server from .conftest import normalize # type: ignore[import-not-found] @@ -178,3 +182,37 @@ def test_delete_assistant_no_router_error(client, auth_as, no_router, path): response = client.delete(path) assert response.status_code == 500 assert len(response.content) > 0 + + +@pytest.mark.parametrize( + ("method", "path", "router_method"), + [ + ("GET", "/v1/assistants", "aget_assistants"), + ("POST", "/v1/assistants", "acreate_assistants"), + ("DELETE", "/v1/assistants/asst_1", "adelete_assistant"), + ], +) +def test_assistants_routes_keep_the_mapped_provider_status( + client: TestClient, + auth_as: Callable[..., AbstractContextManager[None]], + patched_assistants: MagicMock, + method: str, + path: str, + router_method: str, +) -> None: + """A litellm NotFoundError carries status 404 and the SDK's ``code=None``; the route answers 404 with its message.""" + upstream_error: Final = litellm.NotFoundError( + message="NotFoundError: OpenAIException - Error code: 404", + model="gpt-5.4-mini", + llm_provider="openai", + ) + getattr(patched_assistants, router_method).side_effect = upstream_error + with auth_as(): + response: Final = client.request(method, path, json={"model": "gpt-5.4-mini"} if method == "POST" else None) + assert response.status_code == 404 + assert response.json()["error"] == { + "message": upstream_error.message, + "type": "invalid_request_error", + "param": None, + "code": "404", + } diff --git a/tests/unit/proxy/proxy_server/test_routes_threads.py b/tests/unit/proxy/proxy_server/test_routes_threads.py index 493315f041d..bf7bd060694 100644 --- a/tests/unit/proxy/proxy_server/test_routes_threads.py +++ b/tests/unit/proxy/proxy_server/test_routes_threads.py @@ -15,10 +15,14 @@ Pins (PR2): from __future__ import annotations +from contextlib import AbstractContextManager +from typing import Callable, Final from unittest.mock import AsyncMock, MagicMock import pytest +from fastapi.testclient import TestClient +import litellm from litellm.proxy import proxy_server from .conftest import normalize # type: ignore[import-not-found] @@ -272,3 +276,40 @@ def test_run_thread_error(client, auth_as, no_router, path): response = client.post(path, json={"assistant_id": "asst_1"}) assert response.status_code == 500 assert len(response.content) > 0 + + +@pytest.mark.parametrize( + ("method", "path", "router_method", "payload"), + [ + ("POST", "/v1/threads", "acreate_thread", {}), + ("GET", "/v1/threads/thr_1", "aget_thread", None), + ("POST", "/v1/threads/thr_1/messages", "a_add_message", {"role": "user", "content": "hi"}), + ("GET", "/v1/threads/thr_1/messages", "aget_messages", None), + ("POST", "/v1/threads/thr_1/runs", "arun_thread", {"assistant_id": "asst_1"}), + ], +) +def test_threads_routes_keep_the_mapped_provider_status( + client: TestClient, + auth_as: Callable[..., AbstractContextManager[None]], + patched_threads: MagicMock, + method: str, + path: str, + router_method: str, + payload: dict[str, str] | None, +) -> None: + """A litellm NotFoundError carries status 404 and the SDK's ``code=None``; the route answers 404 with its message.""" + upstream_error: Final = litellm.NotFoundError( + message="NotFoundError: OpenAIException - Error code: 404", + model="gpt-5.4-mini", + llm_provider="openai", + ) + getattr(patched_threads, router_method).side_effect = upstream_error + with auth_as(): + response: Final = client.request(method, path, json=payload) + assert response.status_code == 404 + assert response.json()["error"] == { + "message": upstream_error.message, + "type": "invalid_request_error", + "param": None, + "code": "404", + } diff --git a/tests/unit/proxy/test_common_request_processing.py b/tests/unit/proxy/test_common_request_processing.py index cc7cce937d7..d2065d34bd3 100644 --- a/tests/unit/proxy/test_common_request_processing.py +++ b/tests/unit/proxy/test_common_request_processing.py @@ -2346,6 +2346,55 @@ class TestCommonRequestProcessingHelpers: assert plain.code == "429" assert plain.provider_specific_fields is None + async def test_proxy_exception_from_route_error_helper(self) -> None: + """The shared route error -> ProxyException conversion keeps the status an exception + carries, never the OpenAI SDK's ``code`` field (``None`` on most litellm exceptions).""" + from litellm.proxy.common_request_processing import ( + proxy_exception_from_route_error, + ) + + own: Final = ProxyException(message="already shaped", type="invalid_request_error", param=None, code=429) + assert proxy_exception_from_route_error(own) is own + + http: Final = proxy_exception_from_route_error(HTTPException(status_code=403, detail="forbidden")) + assert (http.code, http.type, http.message) == ("403", "permission_error", "forbidden") + + not_found: Final = litellm.NotFoundError( + message="no such assistant", model="gpt-5.4-mini", llm_provider="openai", num_retries=2 + ) + not_found.provider_specific_fields = {"request_id": "req_123"} + assert not_found.code is None + assert str(not_found) != not_found.message + mapped: Final = proxy_exception_from_route_error(not_found) + assert (mapped.code, mapped.type, mapped.param) == ("404", "invalid_request_error", None) + assert mapped.message == not_found.message + assert mapped.provider_specific_fields == {"request_id": "req_123"} + + plain: Final = proxy_exception_from_route_error(ValueError("boom")) + assert (plain.code, plain.type, plain.message) == ("500", "internal_server_error", "boom") + + async def test_proxy_exception_from_route_error_redacts_internal_details(self) -> None: + from litellm.proxy.common_request_processing import ( + proxy_exception_from_route_error, + ) + + notice: Final = bug_report_notice(build_bug_report(RuntimeError("boom"), surface="sdk")) + exc: Final = litellm.APIConnectionError( + message=( + "OpenAIException - postgresql://litellm_internal:S3cr3tPGPass@10.20.30.40:5432/litellm_prod " + f"(config file /etc/litellm/secrets/db.yaml)\n{notice}" + ), + model="gpt-5.4-mini", + llm_provider="openai", + ) + mapped: Final = proxy_exception_from_route_error(exc) + assert mapped.code == "500" + assert "OpenAIException" in mapped.message + assert "REDACTED" in mapped.message + leaked: Final = ("S3cr3tPGPass", "litellm_internal", "10.20.30.40", "/etc/litellm/secrets/db.yaml") + assert [value for value in leaked if value in mapped.message] == [] + assert ISSUE_URL_BASE not in mapped.message + async def test_create_streaming_response_first_chunk_error_string_code(self): """ Test that when the first chunk contains a string error code, a JSON error response is returned