mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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>
This commit is contained in:
parent
2fcc780498
commit
141dba2372
8 changed files with 1134 additions and 121 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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 ####
|
||||
|
|
|
|||
440
tests/integration/providers/_assistants_route_errors_support.py
Normal file
440
tests/integration/providers/_assistants_route_errors_support.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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
|
||||
445
tests/integration/providers/test_assistants_route_errors_wire.py
Normal file
445
tests/integration/providers/test_assistants_route_errors_wire.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue