mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(memory): preserve stream errors and gateway navigation
This commit is contained in:
parent
f79a413e74
commit
5f61a5110b
7 changed files with 71 additions and 18 deletions
|
|
@ -1 +1 @@
|
|||
ALTER TABLE "LiteLLM_MemoryPolicy" ADD COLUMN "paused" BOOLEAN NOT NULL DEFAULT false;
|
||||
ALTER TABLE "LiteLLM_MemoryPolicy" ADD COLUMN IF NOT EXISTS "paused" BOOLEAN NOT NULL DEFAULT false;
|
||||
|
|
|
|||
|
|
@ -27,6 +27,26 @@ def sse_bytes(value: Mapping[str, object], event: str | None = None) -> bytes:
|
|||
return ((f"event: {event}\n" if event else "") + "data: " + json.dumps(value) + "\n\n").encode()
|
||||
|
||||
|
||||
class ServerToolStreamError(ValueError):
|
||||
def __init__(self, data: Mapping[str, object]) -> None:
|
||||
error: Final = object_value(data.get("error") or object_value(data.get("response")).get("error") or data)
|
||||
code: Final = str(error.get("code", ""))
|
||||
self.status_code = (
|
||||
int(code)
|
||||
if code.isascii() and code.isdecimal() and len(code) == 3 and 400 <= int(code) < 600
|
||||
else 429
|
||||
if (error.get("type") or code) in ("rate_limit_error", "rate_limit_exceeded", "insufficient_quota")
|
||||
else 503
|
||||
if error.get("type") == "overloaded_error"
|
||||
else 502
|
||||
)
|
||||
super().__init__(
|
||||
"The upstream model reached its rate or capacity limit"
|
||||
if self.status_code == 429
|
||||
else "The authenticated gateway model stream failed"
|
||||
)
|
||||
|
||||
|
||||
class ServerToolStream:
|
||||
def __init__(
|
||||
self, route: ServerToolRoute, server_names: frozenset[str], request: Mapping[str, object] | None = None
|
||||
|
|
@ -114,7 +134,7 @@ class ServerToolStream:
|
|||
self.frames.append(frame)
|
||||
self.objects.append(data)
|
||||
if data.get("error") or data.get("type") in ("error", "response.failed"):
|
||||
raise ValueError("The model stream failed during gateway tool execution")
|
||||
raise ServerToolStreamError(data)
|
||||
emitted: Final = (
|
||||
self._anthropic(data)
|
||||
if self.route == "anthropic_messages"
|
||||
|
|
@ -424,12 +444,13 @@ class ServerToolStream:
|
|||
b"data: [DONE]\n\n",
|
||||
)
|
||||
|
||||
def error(self, message: str) -> bytes:
|
||||
def error(self, message: str, status_code: int = 502) -> bytes:
|
||||
code: Final = "rate_limit_exceeded" if status_code == 429 else "server_error"
|
||||
if self.route == "aresponses":
|
||||
return self._emit(
|
||||
{ # mutable-ok: Native provider JSON containers.
|
||||
"type": "error",
|
||||
"code": "server_error",
|
||||
"code": code,
|
||||
"message": message,
|
||||
"param": None,
|
||||
},
|
||||
|
|
@ -440,7 +461,7 @@ class ServerToolStream:
|
|||
{ # mutable-ok: Native provider JSON containers.
|
||||
"type": "error",
|
||||
"error": { # mutable-ok: Native provider JSON containers.
|
||||
"type": "api_error",
|
||||
"type": "rate_limit_error" if status_code == 429 else "api_error",
|
||||
"message": message,
|
||||
},
|
||||
},
|
||||
|
|
@ -449,9 +470,9 @@ class ServerToolStream:
|
|||
return sse_bytes(
|
||||
{ # mutable-ok: Native provider JSON containers.
|
||||
"error": { # mutable-ok: Native provider JSON containers.
|
||||
"type": "server_error",
|
||||
"type": code,
|
||||
"message": message,
|
||||
"code": "server_error",
|
||||
"code": str(status_code),
|
||||
}
|
||||
}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -18,7 +18,7 @@ from litellm.litellm_core_utils.prompt_templates.server_tool_responses import (
|
|||
response_has_client_tools,
|
||||
response_messages,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.server_tool_stream import ServerToolStream
|
||||
from litellm.litellm_core_utils.prompt_templates.server_tool_stream import ServerToolStream, ServerToolStreamError
|
||||
from litellm.litellm_core_utils.prompt_templates.server_tools import (
|
||||
ServerToolRoute,
|
||||
append_server_reference,
|
||||
|
|
@ -387,6 +387,8 @@ async def process_gateway_memory(
|
|||
first: Final = await anext(iterator)
|
||||
except StopAsyncIteration as exc:
|
||||
raise HTTPException(status_code=502, detail="The gateway memory stream was empty") from exc
|
||||
except ServerToolStreamError as exc:
|
||||
raise HTTPException(status_code=exc.status_code, detail=str(exc)) from exc
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=502, detail="The gateway model stream was invalid or incomplete") from exc
|
||||
|
||||
|
|
@ -396,8 +398,15 @@ async def process_gateway_memory(
|
|||
async for chunk in iterator:
|
||||
yield chunk
|
||||
except Exception as exc:
|
||||
message: Final = str(exc.detail) if isinstance(exc, HTTPException) else "Gateway memory execution failed"
|
||||
yield loop.stream.error(message)
|
||||
message: Final = (
|
||||
str(exc.detail)
|
||||
if isinstance(exc, HTTPException)
|
||||
else str(exc)
|
||||
if isinstance(exc, ServerToolStreamError)
|
||||
else "Gateway memory execution failed"
|
||||
)
|
||||
status: Final = exc.status_code if isinstance(exc, (HTTPException, ServerToolStreamError)) else 502
|
||||
yield loop.stream.error(message, status)
|
||||
finally:
|
||||
await iterator.aclose()
|
||||
|
||||
|
|
|
|||
|
|
@ -11,7 +11,8 @@ from litellm.litellm_core_utils.prompt_templates.server_tool_responses import (
|
|||
object_value,
|
||||
response_has_client_tools,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.server_tool_stream import ServerToolStream
|
||||
from litellm.litellm_core_utils.prompt_templates.server_tool_stream import ServerToolStream, ServerToolStreamError
|
||||
from litellm.litellm_core_utils.prompt_templates.server_tools import ServerToolRoute
|
||||
|
||||
_MEMORY: Final = frozenset(("litellm_memory_search",))
|
||||
|
||||
|
|
@ -39,6 +40,28 @@ def _event(stream: ServerToolStream, data: Mapping[str, object]) -> bytes:
|
|||
return b"".join(stream.feed(ServerSentEvent(data=json.dumps(data))))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("route", ("acompletion", "aresponses", "anthropic_messages"))
|
||||
@pytest.mark.parametrize(
|
||||
"error",
|
||||
(
|
||||
{"error": {"code": "429", "message": "private upstream account details"}},
|
||||
{"type": "error", "code": "rate_limit_exceeded", "message": "private upstream account details"},
|
||||
{"type": "error", "error": {"type": "rate_limit_error", "message": "private upstream account details"}},
|
||||
),
|
||||
)
|
||||
def test_upstream_stream_rate_limit_retains_classification_without_provider_details(
|
||||
route: ServerToolRoute, error: Mapping[str, object]
|
||||
) -> None:
|
||||
stream: Final = ServerToolStream(route, _MEMORY)
|
||||
with pytest.raises(ServerToolStreamError) as failure:
|
||||
_event(stream, error)
|
||||
assert failure.value.status_code == 429
|
||||
output: Final = stream.error(str(failure.value), failure.value.status_code)
|
||||
assert b"rate_limit" in output
|
||||
assert b"private upstream" not in output
|
||||
assert b"response.completed" not in output
|
||||
|
||||
|
||||
def test_anthropic_stream_hides_memory_keeps_client_tool_ids_and_streams_text_before_completion() -> None:
|
||||
stream: Final = ServerToolStream("anthropic_messages", _MEMORY)
|
||||
start: Final = _event(
|
||||
|
|
|
|||
|
|
@ -237,12 +237,11 @@ def test_every_app_mount_is_assigned_to_a_component():
|
|||
)
|
||||
|
||||
|
||||
def test_memory_v2_policies_stay_on_backend_while_own_entries_are_available_on_gateway():
|
||||
def test_memory_v2_settings_stay_on_backend_while_entries_are_available_on_gateway():
|
||||
gateway = _component_paths(app.router.routes, GATEWAY_EXACT_PATHS, GATEWAY_PATH_PREFIXES)
|
||||
backend = _component_paths(app.router.routes, BACKEND_EXACT_PATHS, BACKEND_PATH_PREFIXES)
|
||||
for path in ("/v2/memory/policies", "/v2/memory/policies/{policy_id}"):
|
||||
assert path in backend
|
||||
assert path not in gateway
|
||||
for path in ("/v2/memory/status", "/v2/memory/preference", "/v2/memory/entries", "/v2/memory/entries/{memory_id}"):
|
||||
assert "/v2/memory/settings" in backend
|
||||
assert "/v2/memory/settings" not in gateway
|
||||
for path in ("/v2/memory/status", "/v2/memory/entries", "/v2/memory/entries/{memory_id}"):
|
||||
assert path in gateway
|
||||
assert path in backend
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ import { Switch } from "@/components/ui/switch";
|
|||
import { fetchClient } from "@/lib/http/api";
|
||||
import type { components } from "@/lib/http/schema";
|
||||
import { toast } from "@/lib/toast";
|
||||
import { uiHref } from "@/utils/uiHref";
|
||||
import { MemoryUserPicker } from "./MemoryTargetPicker";
|
||||
|
||||
type Settings = components["schemas"]["MemorySettings"];
|
||||
|
|
@ -150,7 +151,7 @@ export function MemoryAdministration({
|
|||
all. To give ordinary members access to their team's memories, allow “Read team memories” in Member
|
||||
Permissions.
|
||||
</p>
|
||||
<Link className="inline-block text-sm underline underline-offset-4" href="/teams">
|
||||
<Link className="inline-block text-sm underline underline-offset-4" href={uiHref("teams")}>
|
||||
Manage team permissions
|
||||
</Link>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -126,7 +126,7 @@ describe("Memory dashboard", () => {
|
|||
await waitFor(() => expect(settings.enabled).toBe(true));
|
||||
expect(settings.everyone).toBe(true);
|
||||
await waitFor(() => expect(screen.queryByText("Unsaved changes")).not.toBeInTheDocument());
|
||||
expect(screen.getByRole("link", { name: "Manage team permissions" })).toHaveAttribute("href", "/teams");
|
||||
expect(screen.getByRole("link", { name: "Manage team permissions" })).toHaveAttribute("href", "/ui/teams");
|
||||
await user.click(screen.getByRole("tab", { name: "Memories" }));
|
||||
expect(await screen.findByText("On · Managed by your admin")).toBeVisible();
|
||||
expect(calls.some(({ path }) => path.includes("policies") || path.includes("preference"))).toBe(false);
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue