mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into claude/funny-lamarr-6de68c
# Conflicts: # litellm-proxy-extras/litellm_proxy_extras/schema.prisma # litellm/proxy/schema.prisma # schema.prisma
This commit is contained in:
commit
58105c1c9c
77 changed files with 4898 additions and 304 deletions
8
.github/workflows/create-release-branch.yml
vendored
8
.github/workflows/create-release-branch.yml
vendored
|
|
@ -4,7 +4,7 @@ on:
|
|||
workflow_dispatch:
|
||||
inputs:
|
||||
tag:
|
||||
description: "Release tag (e.g. v1.83.0-stable) — branch will be named release/<tag>"
|
||||
description: "Release tag (e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, 1.84.0.post1; legacy v1.83.10-stable still accepted) — branch will be named release/<tag>"
|
||||
required: true
|
||||
type: string
|
||||
commit_hash:
|
||||
|
|
@ -14,7 +14,7 @@ on:
|
|||
workflow_call:
|
||||
inputs:
|
||||
tag:
|
||||
description: "Release tag"
|
||||
description: "Release tag (e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, 1.84.0.post1; legacy v1.83.10-stable still accepted)"
|
||||
required: true
|
||||
type: string
|
||||
commit_hash:
|
||||
|
|
@ -40,8 +40,8 @@ jobs:
|
|||
echo "::error::commit_hash must be a full 40-character commit SHA"
|
||||
exit 1
|
||||
fi
|
||||
if ! echo "${TAG}" | grep -qE '^v[0-9]+\.[0-9]+\.[0-9]+'; then
|
||||
echo "::error::tag must start with vX.Y.Z"
|
||||
if ! echo "${TAG}" | grep -qE '^v?[0-9]+\.[0-9]+\.[0-9]+'; then
|
||||
echo "::error::tag must start with X.Y.Z (optional leading v), e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, or v1.83.10-stable"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
|
|
|
|||
13
.github/workflows/create-release.yml
vendored
13
.github/workflows/create-release.yml
vendored
|
|
@ -4,7 +4,7 @@ on:
|
|||
workflow_dispatch:
|
||||
inputs:
|
||||
tag:
|
||||
description: "Release tag (e.g. v1.83.0-stable)"
|
||||
description: "Release tag (e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, 1.84.0.post1; legacy v1.83.10-stable still accepted)"
|
||||
required: true
|
||||
type: string
|
||||
commit_hash:
|
||||
|
|
@ -30,8 +30,8 @@ jobs:
|
|||
echo "::error::commit_hash must be a full 40-character commit SHA"
|
||||
exit 1
|
||||
fi
|
||||
if ! echo "${TAG}" | grep -qE '^v[0-9]+\.[0-9]+\.[0-9]+'; then
|
||||
echo "::error::tag must start with vX.Y.Z"
|
||||
if ! echo "${TAG}" | grep -qE '^v?[0-9]+\.[0-9]+\.[0-9]+'; then
|
||||
echo "::error::tag must start with X.Y.Z (optional leading v), e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, or v1.83.10-stable"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
|
|
@ -45,6 +45,11 @@ jobs:
|
|||
const tag = process.env.TAG;
|
||||
const commitHash = process.env.COMMIT_HASH;
|
||||
|
||||
// Mark RC / dev / nightly / alpha / beta tags as GitHub pre-releases.
|
||||
// PEP 440 post-releases (e.g. `1.84.0.post1`) and legacy `-stable[.patch.N]`
|
||||
// are stable maintenance releases, not pre-releases.
|
||||
const isPrerelease = /(?:rc|nightly|alpha|beta|\.dev)/i.test(tag);
|
||||
|
||||
const cosignSection = [
|
||||
`## Verify Docker Image Signature`,
|
||||
``,
|
||||
|
|
@ -89,7 +94,7 @@ jobs:
|
|||
target_commitish: commitHash,
|
||||
name: tag,
|
||||
owner: context.repo.owner,
|
||||
prerelease: false,
|
||||
prerelease: isPrerelease,
|
||||
repo: context.repo.repo,
|
||||
tag_name: tag,
|
||||
});
|
||||
|
|
|
|||
|
|
@ -0,0 +1,75 @@
|
|||
-- CreateTable
|
||||
CREATE TABLE "LiteLLM_WorkflowRun" (
|
||||
"run_id" TEXT NOT NULL,
|
||||
"session_id" TEXT NOT NULL,
|
||||
"workflow_type" TEXT NOT NULL,
|
||||
"status" TEXT NOT NULL DEFAULT 'pending',
|
||||
"created_by" TEXT,
|
||||
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
"updated_at" TIMESTAMP(3) NOT NULL,
|
||||
"input" JSONB,
|
||||
"output" JSONB,
|
||||
"metadata" JSONB,
|
||||
|
||||
CONSTRAINT "LiteLLM_WorkflowRun_pkey" PRIMARY KEY ("run_id")
|
||||
);
|
||||
|
||||
-- CreateTable
|
||||
CREATE TABLE "LiteLLM_WorkflowEvent" (
|
||||
"event_id" TEXT NOT NULL,
|
||||
"run_id" TEXT NOT NULL,
|
||||
"event_type" TEXT NOT NULL,
|
||||
"step_name" TEXT NOT NULL,
|
||||
"sequence_number" INTEGER NOT NULL,
|
||||
"data" JSONB,
|
||||
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
|
||||
CONSTRAINT "LiteLLM_WorkflowEvent_pkey" PRIMARY KEY ("event_id")
|
||||
);
|
||||
|
||||
-- CreateTable
|
||||
CREATE TABLE "LiteLLM_WorkflowMessage" (
|
||||
"message_id" TEXT NOT NULL,
|
||||
"run_id" TEXT NOT NULL,
|
||||
"role" TEXT NOT NULL,
|
||||
"content" TEXT NOT NULL,
|
||||
"sequence_number" INTEGER NOT NULL,
|
||||
"session_id" TEXT,
|
||||
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
|
||||
CONSTRAINT "LiteLLM_WorkflowMessage_pkey" PRIMARY KEY ("message_id")
|
||||
);
|
||||
|
||||
-- CreateIndex
|
||||
CREATE UNIQUE INDEX "LiteLLM_WorkflowRun_session_id_key" ON "LiteLLM_WorkflowRun"("session_id");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE INDEX "LiteLLM_WorkflowRun_workflow_type_status_idx" ON "LiteLLM_WorkflowRun"("workflow_type", "status");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE INDEX "LiteLLM_WorkflowRun_session_id_idx" ON "LiteLLM_WorkflowRun"("session_id");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE INDEX "LiteLLM_WorkflowRun_created_at_idx" ON "LiteLLM_WorkflowRun"("created_at");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE INDEX "LiteLLM_WorkflowRun_created_by_idx" ON "LiteLLM_WorkflowRun"("created_by");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE INDEX "LiteLLM_WorkflowEvent_run_id_idx" ON "LiteLLM_WorkflowEvent"("run_id");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE UNIQUE INDEX "LiteLLM_WorkflowEvent_run_id_sequence_number_key" ON "LiteLLM_WorkflowEvent"("run_id", "sequence_number");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE INDEX "LiteLLM_WorkflowMessage_run_id_idx" ON "LiteLLM_WorkflowMessage"("run_id");
|
||||
|
||||
-- CreateIndex
|
||||
CREATE UNIQUE INDEX "LiteLLM_WorkflowMessage_run_id_sequence_number_key" ON "LiteLLM_WorkflowMessage"("run_id", "sequence_number");
|
||||
|
||||
-- AddForeignKey
|
||||
ALTER TABLE "LiteLLM_WorkflowEvent" ADD CONSTRAINT "LiteLLM_WorkflowEvent_run_id_fkey" FOREIGN KEY ("run_id") REFERENCES "LiteLLM_WorkflowRun"("run_id") ON DELETE RESTRICT ON UPDATE CASCADE;
|
||||
|
||||
-- AddForeignKey
|
||||
ALTER TABLE "LiteLLM_WorkflowMessage" ADD CONSTRAINT "LiteLLM_WorkflowMessage_run_id_fkey" FOREIGN KEY ("run_id") REFERENCES "LiteLLM_WorkflowRun"("run_id") ON DELETE RESTRICT ON UPDATE CASCADE;
|
||||
|
||||
|
|
@ -1291,7 +1291,6 @@ model LiteLLM_AdaptiveRouterSession {
|
|||
@@index([last_activity_at], map: "idx_adaptive_router_session_activity")
|
||||
}
|
||||
|
||||
|
||||
model LiteLLM_ScheduledTaskTable {
|
||||
task_id String @id @default(uuid())
|
||||
// owner_token holds the hashed verification token of the calling key.
|
||||
|
|
@ -1332,3 +1331,80 @@ model LiteLLM_ScheduledTaskTable {
|
|||
@@index([team_id])
|
||||
@@index([agent_id, status])
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Workflow Run Tracking
|
||||
//
|
||||
// Generic durable state tracking for any agent or automated workflow.
|
||||
// Design: three tables — run (header + materialized status), event (append-only
|
||||
// source of truth for state transitions), message (conversation inbox/outbox).
|
||||
//
|
||||
// Usage:
|
||||
// - Set `workflow_type` to identify the owning system (e.g. "shin-builder").
|
||||
// - Store domain-specific fields in `metadata` (worktree_path, pr_url, etc.).
|
||||
// - `session_id` on WorkflowRun matches `x-litellm-session-id` header sent to
|
||||
// the proxy — all spend logs for this run are automatically tagged.
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// One instance of work being done. `status` is a materialized cache of the
|
||||
// latest event; the event log is the authoritative source of truth.
|
||||
model LiteLLM_WorkflowRun {
|
||||
run_id String @id @default(uuid())
|
||||
session_id String @unique @default(uuid())
|
||||
workflow_type String
|
||||
status String @default("pending")
|
||||
created_by String? // user_id of the key that created this run; null = created by master key
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
input Json?
|
||||
output Json?
|
||||
metadata Json?
|
||||
|
||||
events LiteLLM_WorkflowEvent[]
|
||||
messages LiteLLM_WorkflowMessage[]
|
||||
|
||||
@@index([workflow_type, status])
|
||||
@@index([session_id])
|
||||
@@index([created_at])
|
||||
@@index([created_by])
|
||||
}
|
||||
|
||||
// Append-only log of state transitions. Never mutate rows here.
|
||||
// `step_name` and `event_type` are caller-defined strings — no hardcoded enums.
|
||||
// Status auto-update rules (applied by the append endpoint):
|
||||
// step.started → run.status = running
|
||||
// step.failed → run.status = failed
|
||||
// hook.waiting → run.status = paused
|
||||
// hook.received → run.status = running
|
||||
model LiteLLM_WorkflowEvent {
|
||||
event_id String @id @default(uuid())
|
||||
run_id String
|
||||
event_type String
|
||||
step_name String
|
||||
sequence_number Int
|
||||
data Json?
|
||||
created_at DateTime @default(now())
|
||||
|
||||
run LiteLLM_WorkflowRun @relation(fields: [run_id], references: [run_id])
|
||||
|
||||
@@unique([run_id, sequence_number])
|
||||
@@index([run_id])
|
||||
}
|
||||
|
||||
// Conversation inbox/outbox — full message content, separate from the durable
|
||||
// event log. Spend logs truncate messages; this table stores them in full.
|
||||
// `session_id` here is the Claude --resume session ID (or similar).
|
||||
model LiteLLM_WorkflowMessage {
|
||||
message_id String @id @default(uuid())
|
||||
run_id String
|
||||
role String
|
||||
content String
|
||||
sequence_number Int
|
||||
session_id String?
|
||||
created_at DateTime @default(now())
|
||||
|
||||
run LiteLLM_WorkflowRun @relation(fields: [run_id], references: [run_id])
|
||||
|
||||
@@unique([run_id, sequence_number])
|
||||
@@index([run_id])
|
||||
}
|
||||
|
|
|
|||
|
|
@ -224,6 +224,16 @@ AIOHTTP_CONNECTOR_LIMIT_PER_HOST = int(
|
|||
)
|
||||
AIOHTTP_KEEPALIVE_TIMEOUT = int(os.getenv("AIOHTTP_KEEPALIVE_TIMEOUT", 120))
|
||||
AIOHTTP_TTL_DNS_CACHE = int(os.getenv("AIOHTTP_TTL_DNS_CACHE", 300))
|
||||
# TCP keep-alive (SO_KEEPALIVE) — opt-in. Required when running behind NAT/LBs
|
||||
# whose idle timeout is shorter than provider response timeouts (e.g. AWS NAT
|
||||
# Gateway: 350s vs OpenAI/Azure: 600s). Without this, the kernel sends nothing
|
||||
# during a long provider call and the NAT reaps the flow before the response
|
||||
# arrives. Enabling SO_KEEPALIVE makes the kernel emit TCP probes that reset
|
||||
# the NAT idle timer.
|
||||
AIOHTTP_SO_KEEPALIVE = os.getenv("AIOHTTP_SO_KEEPALIVE", "False").lower() == "true"
|
||||
AIOHTTP_TCP_KEEPIDLE = int(os.getenv("AIOHTTP_TCP_KEEPIDLE", 60))
|
||||
AIOHTTP_TCP_KEEPINTVL = int(os.getenv("AIOHTTP_TCP_KEEPINTVL", 30))
|
||||
AIOHTTP_TCP_KEEPCNT = int(os.getenv("AIOHTTP_TCP_KEEPCNT", 5))
|
||||
# enable_cleanup_closed is only needed for Python versions with the SSL leak bug
|
||||
# Fixed in Python 3.12.7+ and 3.13.1+ (see https://github.com/python/cpython/pull/118960)
|
||||
# Reference: https://github.com/aio-libs/aiohttp/blob/master/aiohttp/connector.py#L74-L78
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import asyncio
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, List, Optional
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.proxy.pass_through_endpoints.success_handler import (
|
||||
|
|
@ -29,12 +29,14 @@ class BaseGoogleGenAIGenerateContentStreamingIterator:
|
|||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
request_body: dict,
|
||||
model: str,
|
||||
hidden_params: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
self.litellm_logging_obj = litellm_logging_obj
|
||||
self.request_body = request_body
|
||||
self.start_time = datetime.now()
|
||||
self.collected_chunks: List[bytes] = []
|
||||
self.model = model
|
||||
self._hidden_params: Dict[str, Any] = hidden_params or {}
|
||||
|
||||
async def _handle_async_streaming_logging(
|
||||
self,
|
||||
|
|
@ -76,11 +78,13 @@ class GoogleGenAIGenerateContentStreamingIterator(
|
|||
litellm_metadata: dict,
|
||||
custom_llm_provider: str,
|
||||
request_body: Optional[dict] = None,
|
||||
hidden_params: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
super().__init__(
|
||||
litellm_logging_obj=logging_obj,
|
||||
request_body=request_body or {},
|
||||
model=model,
|
||||
hidden_params=hidden_params,
|
||||
)
|
||||
self.response = response
|
||||
self.model = model
|
||||
|
|
@ -130,11 +134,13 @@ class AsyncGoogleGenAIGenerateContentStreamingIterator(
|
|||
litellm_metadata: dict,
|
||||
custom_llm_provider: str,
|
||||
request_body: Optional[dict] = None,
|
||||
hidden_params: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
super().__init__(
|
||||
litellm_logging_obj=logging_obj,
|
||||
request_body=request_body or {},
|
||||
model=model,
|
||||
hidden_params=hidden_params,
|
||||
)
|
||||
self.response = response
|
||||
self.model = model
|
||||
|
|
|
|||
|
|
@ -348,6 +348,7 @@ def get_llm_provider( # noqa: PLR0915
|
|||
or "ft:gpt-3.5-turbo" in model
|
||||
or "ft:gpt-4" in model # catches ft:gpt-4-0613, ft:gpt-4o
|
||||
or model in litellm.openai_image_generation_models
|
||||
or model.startswith("gpt-image")
|
||||
or model in litellm.openai_video_generation_models
|
||||
):
|
||||
custom_llm_provider = "openai"
|
||||
|
|
|
|||
|
|
@ -982,9 +982,9 @@ class CostCalculatorUtils:
|
|||
image_response=completion_response,
|
||||
)
|
||||
elif custom_llm_provider == litellm.LlmProviders.OPENAI.value:
|
||||
# Check if this is a gpt-image model (token-based pricing)
|
||||
# gpt-image models use token-based pricing.
|
||||
model_lower = model.lower()
|
||||
if "gpt-image-1" in model_lower:
|
||||
if "gpt-image" in model_lower:
|
||||
from litellm.llms.openai.image_generation.cost_calculator import (
|
||||
cost_calculator as openai_gpt_image_cost_calculator,
|
||||
)
|
||||
|
|
@ -1004,9 +1004,9 @@ class CostCalculatorUtils:
|
|||
optional_params=optional_params,
|
||||
)
|
||||
elif custom_llm_provider == litellm.LlmProviders.AZURE.value:
|
||||
# Check if this is a gpt-image model (token-based pricing)
|
||||
# gpt-image models use token-based pricing.
|
||||
model_lower = model.lower()
|
||||
if "gpt-image-1" in model_lower:
|
||||
if "gpt-image" in model_lower:
|
||||
from litellm.llms.openai.image_generation.cost_calculator import (
|
||||
cost_calculator as openai_gpt_image_cost_calculator,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -24,6 +24,6 @@ def get_azure_image_generation_config(model: str) -> BaseImageGenerationConfig:
|
|||
return AzureDallE3ImageGenerationConfig()
|
||||
else:
|
||||
verbose_logger.debug(
|
||||
f"Using AzureGPTImageGenerationConfig for model: {model}. This follows the gpt-image-1 model format."
|
||||
f"Using AzureGPTImageGenerationConfig for model: {model}. This follows the gpt-image model format."
|
||||
)
|
||||
return AzureGPTImageGenerationConfig()
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ from litellm.llms.openai.image_generation import GPTImageGenerationConfig
|
|||
|
||||
class AzureGPTImageGenerationConfig(GPTImageGenerationConfig):
|
||||
"""
|
||||
Azure gpt-image-1 image generation config
|
||||
Azure gpt-image image generation config
|
||||
"""
|
||||
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -95,6 +95,7 @@ class AzureAIVectorStoreConfig(BaseVectorStoreConfig, BaseAzureLLM):
|
|||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[str, Dict[str, Any]]:
|
||||
"""
|
||||
Transform search request for Azure AI Search API
|
||||
|
|
|
|||
|
|
@ -59,6 +59,7 @@ class BaseVectorStoreConfig:
|
|||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[str, Dict]:
|
||||
pass
|
||||
|
||||
|
|
@ -70,6 +71,7 @@ class BaseVectorStoreConfig:
|
|||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[str, Dict]:
|
||||
"""
|
||||
Optional async version of transform_search_vector_store_request.
|
||||
|
|
@ -84,6 +86,7 @@ class BaseVectorStoreConfig:
|
|||
api_base=api_base,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
litellm_params=litellm_params,
|
||||
extra_body=extra_body,
|
||||
)
|
||||
|
||||
@abstractmethod
|
||||
|
|
|
|||
|
|
@ -1942,8 +1942,8 @@ class AmazonConverseConfig(BaseConfig):
|
|||
completion_response = ConverseResponseBlock(**response.json()) # type: ignore
|
||||
except Exception as e:
|
||||
raise BedrockError(
|
||||
message="Received={}, Error converting to valid response block={}. File an issue if litellm error - https://github.com/BerriAI/litellm/issues".format(
|
||||
response.text, str(e)
|
||||
message="Error converting to valid response block={}. File an issue if litellm error - https://github.com/BerriAI/litellm/issues".format(
|
||||
str(e)
|
||||
),
|
||||
status_code=422,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,14 +1,16 @@
|
|||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
|
||||
from copy import deepcopy
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
from litellm.types.integrations.rag.bedrock_knowledgebase import (
|
||||
BedrockKBContent,
|
||||
BedrockKBResponse,
|
||||
BedrockKBRetrievalConfiguration,
|
||||
BedrockKBResponse,
|
||||
BedrockKBRetrievalQuery,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
|
@ -202,6 +204,7 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
|
|||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[str, Dict]:
|
||||
if isinstance(query, list):
|
||||
query = " ".join(query)
|
||||
|
|
@ -213,24 +216,46 @@ class BedrockVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
|
|||
}
|
||||
|
||||
retrieval_config: Dict[str, Any] = {}
|
||||
|
||||
if isinstance(extra_body, dict):
|
||||
retrieval_config = deepcopy(
|
||||
extra_body.get("retrievalConfiguration")
|
||||
or extra_body.get("retrieval_configuration")
|
||||
or {}
|
||||
)
|
||||
max_results = vector_store_search_optional_params.get("max_num_results")
|
||||
if max_results is not None:
|
||||
existing_number_of_results = retrieval_config.get(
|
||||
"vectorSearchConfiguration", {}
|
||||
).get("numberOfResults")
|
||||
if (
|
||||
existing_number_of_results is not None
|
||||
and existing_number_of_results != max_results
|
||||
):
|
||||
verbose_logger.debug(
|
||||
"Overriding extra_body retrievalConfiguration.vectorSearchConfiguration.numberOfResults (%s) with max_num_results=%s",
|
||||
existing_number_of_results,
|
||||
max_results,
|
||||
)
|
||||
retrieval_config.setdefault("vectorSearchConfiguration", {})[
|
||||
"numberOfResults"
|
||||
] = max_results
|
||||
filters = vector_store_search_optional_params.get("filters")
|
||||
if filters is not None:
|
||||
existing_filter = retrieval_config.get("vectorSearchConfiguration", {}).get(
|
||||
"filter"
|
||||
)
|
||||
if existing_filter is not None and existing_filter != filters:
|
||||
verbose_logger.debug(
|
||||
"Overriding extra_body retrievalConfiguration.vectorSearchConfiguration.filter with filters from vector_store_search_optional_params"
|
||||
)
|
||||
retrieval_config.setdefault("vectorSearchConfiguration", {})[
|
||||
"filter"
|
||||
] = filters
|
||||
if retrieval_config:
|
||||
# Create a properly typed retrieval configuration
|
||||
typed_retrieval_config: BedrockKBRetrievalConfiguration = {}
|
||||
if "vectorSearchConfiguration" in retrieval_config:
|
||||
typed_retrieval_config["vectorSearchConfiguration"] = retrieval_config[
|
||||
"vectorSearchConfiguration"
|
||||
]
|
||||
request_body["retrievalConfiguration"] = typed_retrieval_config
|
||||
request_body["retrievalConfiguration"] = cast(
|
||||
BedrockKBRetrievalConfiguration, retrieval_config
|
||||
)
|
||||
|
||||
litellm_logging_obj.model_call_details["query"] = query
|
||||
return url, request_body
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
import asyncio
|
||||
import inspect
|
||||
import os
|
||||
import socket
|
||||
import ssl
|
||||
import sys
|
||||
import time
|
||||
|
|
@ -29,6 +31,10 @@ from litellm.constants import (
|
|||
AIOHTTP_CONNECTOR_LIMIT_PER_HOST,
|
||||
AIOHTTP_KEEPALIVE_TIMEOUT,
|
||||
AIOHTTP_NEEDS_CLEANUP_CLOSED,
|
||||
AIOHTTP_SO_KEEPALIVE,
|
||||
AIOHTTP_TCP_KEEPCNT,
|
||||
AIOHTTP_TCP_KEEPIDLE,
|
||||
AIOHTTP_TCP_KEEPINTVL,
|
||||
AIOHTTP_TTL_DNS_CACHE,
|
||||
COMPLETION_HTTP_FALLBACK_SECONDS,
|
||||
DEFAULT_SSL_CIPHERS,
|
||||
|
|
@ -54,6 +60,57 @@ except Exception:
|
|||
version = "0.0.0"
|
||||
|
||||
|
||||
# aiohttp 3.10+ exposes a `socket_factory` kwarg on TCPConnector. Older
|
||||
# versions don't — detect once and skip the keep-alive wiring there.
|
||||
# https://docs.aiohttp.org/en/stable/client_reference.html#aiohttp.TCPConnector
|
||||
_AIOHTTP_SUPPORTS_SOCKET_FACTORY = (
|
||||
"socket_factory" in inspect.signature(TCPConnector.__init__).parameters
|
||||
)
|
||||
|
||||
|
||||
def _build_aiohttp_keepalive_socket_factory() -> (
|
||||
Optional[Callable[[Tuple[Any, ...]], socket.socket]]
|
||||
):
|
||||
"""
|
||||
Build a socket_factory that enables SO_KEEPALIVE on aiohttp TCP sockets.
|
||||
|
||||
Why: by default, aiohttp creates sockets without SO_KEEPALIVE, so the kernel
|
||||
sends nothing during a long idle TCP connection. NAT/LB hops (e.g. AWS NAT
|
||||
Gateway, 350s idle timeout) reap the flow well before slow provider
|
||||
responses (OpenAI/Azure: up to 600s) arrive. Enabling SO_KEEPALIVE makes
|
||||
the kernel emit TCP probes that reset the NAT idle timer.
|
||||
|
||||
Returns None when AIOHTTP_SO_KEEPALIVE is disabled or aiohttp is too old.
|
||||
"""
|
||||
if not AIOHTTP_SO_KEEPALIVE or not _AIOHTTP_SUPPORTS_SOCKET_FACTORY:
|
||||
return None
|
||||
|
||||
def factory(addr_info: Tuple[Any, ...]) -> socket.socket:
|
||||
family, type_, proto = addr_info[0], addr_info[1], addr_info[2]
|
||||
sock = socket.socket(family=family, type=type_, proto=proto)
|
||||
sock.setblocking(False)
|
||||
sock.setsockopt(socket.SOL_SOCKET, socket.SO_KEEPALIVE, 1)
|
||||
# Linux: TCP_KEEPIDLE is idle-before-first-probe.
|
||||
# macOS/Darwin: TCP_KEEPALIVE is the equivalent.
|
||||
if hasattr(socket, "TCP_KEEPIDLE"):
|
||||
sock.setsockopt(
|
||||
socket.IPPROTO_TCP, socket.TCP_KEEPIDLE, AIOHTTP_TCP_KEEPIDLE
|
||||
)
|
||||
elif hasattr(socket, "TCP_KEEPALIVE"):
|
||||
sock.setsockopt(
|
||||
socket.IPPROTO_TCP, socket.TCP_KEEPALIVE, AIOHTTP_TCP_KEEPIDLE
|
||||
)
|
||||
if hasattr(socket, "TCP_KEEPINTVL"):
|
||||
sock.setsockopt(
|
||||
socket.IPPROTO_TCP, socket.TCP_KEEPINTVL, AIOHTTP_TCP_KEEPINTVL
|
||||
)
|
||||
if hasattr(socket, "TCP_KEEPCNT"):
|
||||
sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_KEEPCNT, AIOHTTP_TCP_KEEPCNT)
|
||||
return sock
|
||||
|
||||
return factory
|
||||
|
||||
|
||||
def get_default_headers() -> dict:
|
||||
"""
|
||||
Get default headers for HTTP requests.
|
||||
|
|
@ -935,6 +992,11 @@ class AsyncHTTPHandler:
|
|||
transport_connector_kwargs["limit_per_host"] = (
|
||||
AIOHTTP_CONNECTOR_LIMIT_PER_HOST
|
||||
)
|
||||
# Returns None when SO_KEEPALIVE is disabled or aiohttp is too old to
|
||||
# accept socket_factory — version detection lives inside the builder.
|
||||
socket_factory = _build_aiohttp_keepalive_socket_factory()
|
||||
if socket_factory is not None:
|
||||
transport_connector_kwargs["socket_factory"] = socket_factory
|
||||
|
||||
return LiteLLMAiohttpTransport(
|
||||
client=lambda: ClientSession(
|
||||
|
|
|
|||
|
|
@ -155,6 +155,30 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
def _google_genai_streaming_hidden_params(
|
||||
*,
|
||||
api_base: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
response_headers: httpx.Headers,
|
||||
) -> Dict[str, Any]:
|
||||
"""Pre-stream metadata for proxy response headers (mirrors CustomStreamWrapper._hidden_params)."""
|
||||
from litellm.litellm_core_utils.core_helpers import process_response_headers
|
||||
|
||||
_model_info: Dict[str, Any] = dict(
|
||||
getattr(litellm_params, "model_info", None) or {}
|
||||
)
|
||||
_raw_id = _model_info.get("id") or logging_obj.get_router_model_id() or ""
|
||||
_model_id = _raw_id if isinstance(_raw_id, str) else str(_raw_id)
|
||||
return {
|
||||
"model_id": _model_id,
|
||||
"api_base": api_base,
|
||||
"cache_key": "",
|
||||
"response_cost": "",
|
||||
"additional_headers": process_response_headers(response_headers),
|
||||
}
|
||||
|
||||
|
||||
class BaseLLMHTTPHandler:
|
||||
async def _make_common_async_call(
|
||||
self,
|
||||
|
|
@ -8585,6 +8609,7 @@ class BaseLLMHTTPHandler:
|
|||
api_base=api_base,
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params=dict(litellm_params),
|
||||
extra_body=extra_body,
|
||||
)
|
||||
else:
|
||||
(
|
||||
|
|
@ -8597,6 +8622,7 @@ class BaseLLMHTTPHandler:
|
|||
api_base=api_base,
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params=dict(litellm_params),
|
||||
extra_body=extra_body,
|
||||
)
|
||||
all_optional_params: Dict[str, Any] = dict(litellm_params)
|
||||
all_optional_params.update(vector_store_search_optional_params or {})
|
||||
|
|
@ -8697,6 +8723,7 @@ class BaseLLMHTTPHandler:
|
|||
api_base=api_base,
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params=dict(litellm_params),
|
||||
extra_body=extra_body,
|
||||
)
|
||||
|
||||
all_optional_params: Dict[str, Any] = dict(litellm_params)
|
||||
|
|
@ -10425,6 +10452,12 @@ class BaseLLMHTTPHandler:
|
|||
litellm_metadata=litellm_metadata or {},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
request_body=data,
|
||||
hidden_params=_google_genai_streaming_hidden_params(
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
response_headers=response.headers,
|
||||
),
|
||||
)
|
||||
else:
|
||||
response = sync_httpx_client.post(
|
||||
|
|
@ -10534,6 +10567,12 @@ class BaseLLMHTTPHandler:
|
|||
litellm_metadata=litellm_metadata or {},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
request_body=data,
|
||||
hidden_params=_google_genai_streaming_hidden_params(
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=logging_obj,
|
||||
response_headers=response.headers,
|
||||
),
|
||||
)
|
||||
else:
|
||||
response = await async_httpx_client.post(
|
||||
|
|
|
|||
|
|
@ -118,6 +118,7 @@ class GeminiVectorStoreConfig(BaseVectorStoreConfig):
|
|||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[str, Dict]:
|
||||
"""
|
||||
Transform search request to Gemini's generateContent format.
|
||||
|
|
|
|||
|
|
@ -130,6 +130,7 @@ class MilvusVectorStoreConfig(BaseVectorStoreConfig):
|
|||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[str, Dict[str, Any]]:
|
||||
"""
|
||||
Transform search request for Azure AI Search API
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
"""
|
||||
Cost calculator for OpenAI image generation models (gpt-image-1, gpt-image-1-mini)
|
||||
Cost calculator for OpenAI image generation models (gpt-image family)
|
||||
|
||||
These models use token-based pricing instead of pixel-based pricing like DALL-E.
|
||||
"""
|
||||
|
|
@ -17,13 +17,13 @@ def cost_calculator(
|
|||
custom_llm_provider: Optional[str] = None,
|
||||
) -> float:
|
||||
"""
|
||||
Calculate cost for OpenAI gpt-image-1 and gpt-image-1-mini models.
|
||||
Calculate cost for OpenAI gpt-image models.
|
||||
|
||||
Uses the same usage format as Responses API, so we reuse the helper
|
||||
to transform to chat completion format and use generic_cost_per_token.
|
||||
|
||||
Args:
|
||||
model: The model name (e.g., "gpt-image-1", "gpt-image-1-mini")
|
||||
model: The model name (e.g., "gpt-image-1", "gpt-image-2")
|
||||
image_response: The ImageResponse containing usage data
|
||||
custom_llm_provider: Optional provider name
|
||||
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@ if TYPE_CHECKING:
|
|||
|
||||
class GPTImageGenerationConfig(BaseImageGenerationConfig):
|
||||
"""
|
||||
OpenAI gpt-image-1 image generation config
|
||||
OpenAI gpt-image image generation config
|
||||
"""
|
||||
|
||||
def get_supported_openai_params(
|
||||
|
|
|
|||
|
|
@ -106,6 +106,7 @@ class OpenAIVectorStoreConfig(BaseVectorStoreConfig):
|
|||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[str, Dict]:
|
||||
url = f"{api_base}/{vector_store_id}/search"
|
||||
typed_request_body = VectorStoreSearchRequest(
|
||||
|
|
|
|||
|
|
@ -101,5 +101,10 @@
|
|||
"param_mappings": {
|
||||
"max_completion_tokens": "max_tokens"
|
||||
}
|
||||
},
|
||||
"aihubmix": {
|
||||
"base_url": "https://aihubmix.com/v1",
|
||||
"api_key_env": "AIHUBMIX_API_KEY",
|
||||
"api_base_env": "AIHUBMIX_API_BASE"
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -80,6 +80,7 @@ class PGVectorStoreConfig(OpenAIVectorStoreConfig):
|
|||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[str, Dict]:
|
||||
url = f"{api_base}/{vector_store_id}/search"
|
||||
_, request_body = super().transform_search_vector_store_request(
|
||||
|
|
@ -89,5 +90,6 @@ class PGVectorStoreConfig(OpenAIVectorStoreConfig):
|
|||
api_base=api_base,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
litellm_params=litellm_params,
|
||||
extra_body=extra_body,
|
||||
)
|
||||
return url, request_body
|
||||
|
|
|
|||
|
|
@ -102,6 +102,7 @@ class RAGFlowVectorStoreConfig(BaseVectorStoreConfig):
|
|||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[str, Dict]:
|
||||
"""RAGFlow vector stores are management-only, search is not supported."""
|
||||
raise NotImplementedError(
|
||||
|
|
|
|||
|
|
@ -79,6 +79,7 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
|
|||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[str, Dict]:
|
||||
"""Sync version - generates embedding synchronously."""
|
||||
# For S3 Vectors, vector_store_id should be in format: bucket_name:index_name
|
||||
|
|
@ -140,6 +141,7 @@ class S3VectorsVectorStoreConfig(BaseVectorStoreConfig, BaseAWSLLM):
|
|||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[str, Dict]:
|
||||
"""Async version - generates embedding asynchronously."""
|
||||
# For S3 Vectors, vector_store_id should be in format: bucket_name:index_name
|
||||
|
|
|
|||
|
|
@ -2395,8 +2395,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
completion_response = GenerateContentResponseBody(**raw_response.json()) # type: ignore
|
||||
except Exception as e:
|
||||
raise VertexAIError(
|
||||
message="Received={}, Error converting to valid response block={}. File an issue if litellm error - https://github.com/BerriAI/litellm/issues".format(
|
||||
raw_response.text, str(e)
|
||||
message="Error converting to valid response block={}. File an issue if litellm error - https://github.com/BerriAI/litellm/issues".format(
|
||||
str(e)
|
||||
),
|
||||
status_code=422,
|
||||
headers=raw_response.headers,
|
||||
|
|
@ -2530,8 +2530,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
except Exception as e:
|
||||
raise VertexAIError(
|
||||
message="Received={}, Error converting to valid response block={}. File an issue if litellm error - https://github.com/BerriAI/litellm/issues".format(
|
||||
completion_response, str(e)
|
||||
message="Error converting to valid response block={}. File an issue if litellm error - https://github.com/BerriAI/litellm/issues".format(
|
||||
str(e)
|
||||
),
|
||||
status_code=422,
|
||||
headers=raw_response.headers,
|
||||
|
|
|
|||
|
|
@ -100,6 +100,7 @@ class VertexVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
|
|||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[str, Dict[str, Any]]:
|
||||
"""
|
||||
Transform search request for Vertex AI RAG API
|
||||
|
|
|
|||
|
|
@ -107,6 +107,7 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
|
|||
api_base: str,
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
extra_body: Optional[Dict[str, Any]] = None,
|
||||
) -> Tuple[str, Dict[str, Any]]:
|
||||
"""
|
||||
Transform search request for Vertex AI RAG API
|
||||
|
|
|
|||
|
|
@ -5103,6 +5103,38 @@
|
|||
"/v1/images/edits"
|
||||
]
|
||||
},
|
||||
"azure/gpt-image-2": {
|
||||
"cache_read_input_image_token_cost": 2e-06,
|
||||
"cache_read_input_token_cost": 1.25e-06,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"input_cost_per_image_token": 8e-06,
|
||||
"litellm_provider": "azure",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_token": 1e-05,
|
||||
"output_cost_per_image_token": 3e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supports_vision": true,
|
||||
"supports_pdf_input": true
|
||||
},
|
||||
"azure/gpt-image-2-2026-04-21": {
|
||||
"cache_read_input_image_token_cost": 2e-06,
|
||||
"cache_read_input_token_cost": 1.25e-06,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"input_cost_per_image_token": 8e-06,
|
||||
"litellm_provider": "azure",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_token": 1e-05,
|
||||
"output_cost_per_image_token": 3e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supports_vision": true,
|
||||
"supports_pdf_input": true
|
||||
},
|
||||
"azure/low/1024-x-1024/gpt-image-1-mini": {
|
||||
"input_cost_per_pixel": 2.0751953125e-09,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -19083,6 +19115,38 @@
|
|||
"supports_vision": true,
|
||||
"supports_pdf_input": true
|
||||
},
|
||||
"gpt-image-2": {
|
||||
"cache_read_input_image_token_cost": 2e-06,
|
||||
"cache_read_input_token_cost": 1.25e-06,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_token": 1e-05,
|
||||
"input_cost_per_image_token": 8e-06,
|
||||
"output_cost_per_image_token": 3e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supports_vision": true,
|
||||
"supports_pdf_input": true
|
||||
},
|
||||
"gpt-image-2-2026-04-21": {
|
||||
"cache_read_input_image_token_cost": 2e-06,
|
||||
"cache_read_input_token_cost": 1.25e-06,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_token": 1e-05,
|
||||
"input_cost_per_image_token": 8e-06,
|
||||
"output_cost_per_image_token": 3e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supports_vision": true,
|
||||
"supports_pdf_input": true
|
||||
},
|
||||
"low/1024-x-1024/gpt-image-1.5": {
|
||||
"input_cost_per_image": 0.009,
|
||||
"litellm_provider": "openai",
|
||||
|
|
|
|||
|
|
@ -50,8 +50,11 @@ from litellm.proxy._experimental.mcp_server.oauth2_token_cache import resolve_mc
|
|||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
MCP_TOOL_PREFIX_SEPARATOR,
|
||||
add_server_prefix_to_name,
|
||||
compute_short_server_prefix,
|
||||
get_server_prefix,
|
||||
is_short_mcp_tool_prefix_enabled,
|
||||
is_tool_name_prefixed,
|
||||
iter_known_server_prefixes,
|
||||
merge_mcp_headers,
|
||||
normalize_server_name,
|
||||
split_server_prefix_from_name,
|
||||
|
|
@ -106,6 +109,12 @@ if not _separator_probe.is_valid:
|
|||
SEP_986_URL,
|
||||
)
|
||||
|
||||
_AZURE_ENTRA_HOSTS = {
|
||||
"login.microsoftonline.com", # Global
|
||||
"login.microsoftonline.us", # US Government
|
||||
"login.chinacloudapi.cn", # China
|
||||
}
|
||||
|
||||
|
||||
def _warn_on_server_name_fields(
|
||||
*,
|
||||
|
|
@ -364,6 +373,7 @@ class MCPServerManager:
|
|||
aws_session_name=server_config.get("aws_session_name", None),
|
||||
instructions=server_config.get("instructions", None),
|
||||
)
|
||||
self._assign_unique_short_prefix(new_server)
|
||||
self.config_mcp_servers[server_id] = new_server
|
||||
|
||||
# Check if this is an OpenAPI-based server
|
||||
|
|
@ -726,6 +736,7 @@ class MCPServerManager:
|
|||
try:
|
||||
if mcp_server.server_id not in self.registry:
|
||||
new_server = await self.build_mcp_server_from_table(mcp_server)
|
||||
self._assign_unique_short_prefix(new_server)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
verbose_logger.debug(f"Added MCP Server: {new_server.name}")
|
||||
|
|
@ -738,6 +749,12 @@ class MCPServerManager:
|
|||
try:
|
||||
if mcp_server.server_id in self.registry:
|
||||
new_server = await self.build_mcp_server_from_table(mcp_server)
|
||||
# Carry the previously-resolved short prefix across so the
|
||||
# tool names stay stable for clients holding cached lists.
|
||||
existing_prefix = self.registry[mcp_server.server_id].short_prefix
|
||||
if existing_prefix and not new_server.short_prefix:
|
||||
new_server.short_prefix = existing_prefix
|
||||
self._assign_unique_short_prefix(new_server)
|
||||
self.registry[mcp_server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
verbose_logger.debug(f"Updated MCP Server: {new_server.name}")
|
||||
|
|
@ -1236,7 +1253,11 @@ class MCPServerManager:
|
|||
|
||||
## HANDLE OPENAPI TOOLS
|
||||
if server.spec_path:
|
||||
_tools = global_mcp_tool_registry.list_tools(tool_prefix=server.name)
|
||||
# OpenAPI tools were stored in the registry under the prefix
|
||||
# active at registration time — fetch by that same prefix.
|
||||
_tools = global_mcp_tool_registry.list_tools(
|
||||
tool_prefix=get_server_prefix(server)
|
||||
)
|
||||
tools = global_mcp_tool_registry.convert_tools_to_mcp_sdk_tool_type(
|
||||
_tools
|
||||
)
|
||||
|
|
@ -1488,11 +1509,28 @@ class MCPServerManager:
|
|||
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP)
|
||||
response = await client.get(server_url)
|
||||
response.raise_for_status()
|
||||
verbose_logger.warning(
|
||||
"MCP OAuth discovery unexpectedly succeeded for %s; server did not challenge",
|
||||
server_url,
|
||||
(
|
||||
authorization_servers,
|
||||
resource_scopes,
|
||||
) = await self._attempt_well_known_discovery(server_url)
|
||||
metadata = await self._fetch_authorization_server_metadata(
|
||||
authorization_servers
|
||||
)
|
||||
raise RuntimeError("OAuth discovery must not succeed without a challenge")
|
||||
if (
|
||||
metadata is None
|
||||
and not resource_scopes
|
||||
and authorization_servers
|
||||
and response.status_code == 200
|
||||
):
|
||||
verbose_logger.warning(
|
||||
"MCP OAuth discovery for %s received 200 OK without RFC 9728 challenge and no discoverable authorization metadata.",
|
||||
server_url,
|
||||
)
|
||||
if metadata is None and resource_scopes:
|
||||
return MCPOAuthMetadata(scopes=resource_scopes)
|
||||
if metadata is not None and resource_scopes:
|
||||
metadata.scopes = resource_scopes
|
||||
return metadata
|
||||
except HTTPStatusError as exc:
|
||||
verbose_logger.debug(
|
||||
"MCP OAuth discovery for %s received status error: %s",
|
||||
|
|
@ -1510,8 +1548,8 @@ class MCPServerManager:
|
|||
header_value
|
||||
)
|
||||
|
||||
authorization_servers: List[str] = []
|
||||
resource_scopes: Optional[List[str]] = None
|
||||
authorization_servers = []
|
||||
resource_scopes = None
|
||||
if resource_metadata_url:
|
||||
(
|
||||
authorization_servers,
|
||||
|
|
@ -1674,6 +1712,9 @@ class MCPServerManager:
|
|||
f"{base}/.well-known/oauth-authorization-server/{path}"
|
||||
)
|
||||
candidate_urls.append(f"{base}/.well-known/openid-configuration/{path}")
|
||||
candidate_urls.append(
|
||||
f"{issuer_url.rstrip('/')}/.well-known/openid-configuration"
|
||||
)
|
||||
candidate_urls.append(f"{base}/.well-known/oauth-authorization-server")
|
||||
candidate_urls.append(f"{base}/.well-known/openid-configuration")
|
||||
candidate_urls.append(issuer_url.rstrip("/"))
|
||||
|
|
@ -1713,7 +1754,28 @@ class MCPServerManager:
|
|||
):
|
||||
return metadata
|
||||
|
||||
return None
|
||||
return self._build_azure_authorization_server_metadata(parsed)
|
||||
|
||||
@staticmethod
|
||||
def _build_azure_authorization_server_metadata(
|
||||
parsed_issuer_url: Any,
|
||||
) -> Optional[MCPOAuthMetadata]:
|
||||
path_parts = [
|
||||
part for part in (parsed_issuer_url.path or "").split("/") if part
|
||||
]
|
||||
if (
|
||||
parsed_issuer_url.netloc not in _AZURE_ENTRA_HOSTS
|
||||
or len(path_parts) != 2
|
||||
or path_parts[1] != "v2.0"
|
||||
):
|
||||
return None
|
||||
|
||||
tenant = path_parts[0]
|
||||
base = f"{parsed_issuer_url.scheme}://{parsed_issuer_url.netloc}/{tenant}"
|
||||
return MCPOAuthMetadata(
|
||||
authorization_url=f"{base}/oauth2/v2.0/authorize",
|
||||
token_url=f"{base}/oauth2/v2.0/token",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _decrypt_credential_field(
|
||||
|
|
@ -1810,6 +1872,63 @@ class MCPServerManager:
|
|||
verbose_logger.warning(f"Error listing tools from {server_name}: {str(e)}")
|
||||
return []
|
||||
|
||||
_SHORT_PREFIX_MAX_REHASH_ATTEMPTS = 1024
|
||||
|
||||
def _assign_unique_short_prefix(self, server: MCPServer) -> None:
|
||||
"""Resolve and cache a collision-free short tool prefix on ``server``.
|
||||
|
||||
Called at registration time for every MCP server entering the
|
||||
registry. Mutates ``server.short_prefix`` in place. No-ops when
|
||||
``LITELLM_USE_SHORT_MCP_TOOL_PREFIX`` is disabled, when the server
|
||||
has no ``server_id`` (synthetic temp-server objects), or when a
|
||||
prefix is already cached.
|
||||
|
||||
Collision strategy: take the natural hash; if it's already used by
|
||||
a *different* server in the combined registry, rehash with an
|
||||
incrementing attempt counter until we find an unused slot. The
|
||||
attempt counter is folded into the hash so the resulting prefix is
|
||||
still deterministic for a given (server_id, set-of-other-server-ids)
|
||||
pair within one process.
|
||||
"""
|
||||
if not is_short_mcp_tool_prefix_enabled():
|
||||
return
|
||||
if server.short_prefix:
|
||||
return
|
||||
if not server.server_id:
|
||||
return
|
||||
|
||||
used: Dict[str, str] = {}
|
||||
for other in self.get_registry().values():
|
||||
if other.server_id == server.server_id:
|
||||
continue
|
||||
if other.short_prefix:
|
||||
used[other.short_prefix] = other.server_id
|
||||
|
||||
for attempt in range(self._SHORT_PREFIX_MAX_REHASH_ATTEMPTS):
|
||||
candidate = compute_short_server_prefix(server.server_id, attempt=attempt)
|
||||
if candidate not in used:
|
||||
server.short_prefix = candidate
|
||||
if attempt > 0:
|
||||
verbose_logger.info(
|
||||
"MCP short-prefix collision resolved for server %s: "
|
||||
"natural hash collided with %s, using rehashed prefix "
|
||||
"%s (attempt=%d).",
|
||||
server.server_id,
|
||||
used.get(
|
||||
compute_short_server_prefix(server.server_id, attempt=0),
|
||||
"<unknown>",
|
||||
),
|
||||
candidate,
|
||||
attempt,
|
||||
)
|
||||
return
|
||||
|
||||
raise RuntimeError(
|
||||
f"Unable to assign a unique short MCP tool prefix for server "
|
||||
f"{server.server_id} after {self._SHORT_PREFIX_MAX_REHASH_ATTEMPTS} "
|
||||
"attempts; the 3-character prefix space is too crowded."
|
||||
)
|
||||
|
||||
def _create_prefixed_tools(
|
||||
self, tools: List[MCPTool], server: MCPServer, add_prefix: bool = True
|
||||
) -> List[MCPTool]:
|
||||
|
|
@ -1838,9 +1957,13 @@ class MCPServerManager:
|
|||
tool_copy.name = name_to_use
|
||||
prefixed_tools.append(tool_copy)
|
||||
|
||||
# Update tool to server mapping for resolution (support both forms)
|
||||
# Register every known prefix form (alias, server_name, server_id,
|
||||
# short ID) so call_tool can resolve regardless of which form a
|
||||
# caller / cached client is using.
|
||||
self.tool_name_to_mcp_server_name_mapping[original_name] = prefix
|
||||
self.tool_name_to_mcp_server_name_mapping[prefixed_name] = prefix
|
||||
for known_prefix in iter_known_server_prefixes(server):
|
||||
qualified = add_server_prefix_to_name(original_name, known_prefix)
|
||||
self.tool_name_to_mcp_server_name_mapping[qualified] = prefix
|
||||
|
||||
verbose_logger.info(
|
||||
f"Successfully fetched {len(prefixed_tools)} tools from server {server.name}"
|
||||
|
|
@ -2601,37 +2724,43 @@ class MCPServerManager:
|
|||
Returns:
|
||||
MCPServer if found, None otherwise
|
||||
"""
|
||||
registry_servers = list(self.get_registry().values())
|
||||
|
||||
# Build prefix → server lookup covering every known form a tool name
|
||||
# may take (alias / server_name / server_id / short ID). This is what
|
||||
# makes the short-prefix mode work without breaking historical names.
|
||||
prefix_to_server: Dict[str, MCPServer] = {}
|
||||
for server in registry_servers:
|
||||
for known_prefix in iter_known_server_prefixes(server):
|
||||
normalised = normalize_server_name(known_prefix)
|
||||
prefix_to_server.setdefault(normalised, server)
|
||||
|
||||
# First try with the original tool name
|
||||
if tool_name in self.tool_name_to_mcp_server_name_mapping:
|
||||
server_name = self.tool_name_to_mcp_server_name_mapping[tool_name]
|
||||
for server in self.get_registry().values():
|
||||
if normalize_server_name(server.name) == normalize_server_name(
|
||||
server_name
|
||||
):
|
||||
normalised_lookup = normalize_server_name(server_name)
|
||||
if normalised_lookup in prefix_to_server:
|
||||
return prefix_to_server[normalised_lookup]
|
||||
for server in registry_servers:
|
||||
if normalize_server_name(server.name) == normalised_lookup:
|
||||
return server
|
||||
|
||||
# If not found and tool name is prefixed, try extracting server name from prefix
|
||||
known_prefixes = {
|
||||
normalize_server_name(get_server_prefix(s))
|
||||
for s in self.get_registry().values()
|
||||
if get_server_prefix(s)
|
||||
}
|
||||
if is_tool_name_prefixed(tool_name, known_server_prefixes=known_prefixes):
|
||||
# If not found and tool name is prefixed, extract the prefix and
|
||||
# match against any known form.
|
||||
if is_tool_name_prefixed(
|
||||
tool_name, known_server_prefixes=set(prefix_to_server.keys())
|
||||
):
|
||||
(
|
||||
original_tool_name,
|
||||
server_name_from_prefix,
|
||||
) = split_server_prefix_from_name(tool_name)
|
||||
if original_tool_name in self.tool_name_to_mcp_server_name_mapping:
|
||||
for server in self.get_registry().values():
|
||||
if server.server_name is None:
|
||||
if normalize_server_name(server.name) == normalize_server_name(
|
||||
server_name_from_prefix
|
||||
):
|
||||
return server
|
||||
elif normalize_server_name(
|
||||
server.server_name
|
||||
) == normalize_server_name(server_name_from_prefix):
|
||||
return server
|
||||
normalised_prefix = normalize_server_name(server_name_from_prefix)
|
||||
matched_server = prefix_to_server.get(normalised_prefix)
|
||||
if matched_server is not None and (
|
||||
original_tool_name in self.tool_name_to_mcp_server_name_mapping
|
||||
or tool_name in self.tool_name_to_mcp_server_name_mapping
|
||||
):
|
||||
return matched_server
|
||||
|
||||
return None
|
||||
|
||||
|
|
@ -2666,6 +2795,9 @@ class MCPServerManager:
|
|||
previous_registry = self.registry
|
||||
new_registry: Dict[str, MCPServer] = {}
|
||||
|
||||
# Stage one: build every server. Stage two assigns short prefixes
|
||||
# against the *full* set so dedup is deterministic regardless of
|
||||
# iteration order.
|
||||
for server in db_mcp_servers:
|
||||
existing_server = previous_registry.get(server.server_id)
|
||||
|
||||
|
|
@ -2689,10 +2821,21 @@ class MCPServerManager:
|
|||
f"Building server from DB: {server.server_id} ({server.server_name})"
|
||||
)
|
||||
new_server = await self.build_mcp_server_from_table(server)
|
||||
# Carry the cached short_prefix from the previous registry entry
|
||||
# (if any) so the prefix is stable across reloads.
|
||||
if existing_server is not None and existing_server.short_prefix:
|
||||
new_server.short_prefix = existing_server.short_prefix
|
||||
new_registry[server.server_id] = new_server
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
|
||||
# Swap in the new registry first so _assign_unique_short_prefix
|
||||
# sees the complete set when checking for collisions.
|
||||
self.registry = new_registry
|
||||
for new_server in new_registry.values():
|
||||
self._assign_unique_short_prefix(new_server)
|
||||
# Register OpenAPI tools *after* the final short prefix is assigned
|
||||
# so the tools are stored in the global registry under the same
|
||||
# prefix that lookups will use.
|
||||
await self._maybe_register_openapi_tools(new_server)
|
||||
|
||||
verbose_logger.debug(
|
||||
"MCP registry refreshed (%s servers in registry)", len(new_registry)
|
||||
|
|
|
|||
|
|
@ -49,6 +49,7 @@ from litellm.proxy._experimental.mcp_server.utils import (
|
|||
LITELLM_MCP_SERVER_VERSION,
|
||||
add_server_prefix_to_name,
|
||||
get_server_prefix,
|
||||
iter_known_server_prefixes,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
|
|
@ -711,13 +712,7 @@ if MCP_AVAILABLE:
|
|||
for server in allowed_mcp_servers:
|
||||
if server:
|
||||
match_list = [
|
||||
s.lower()
|
||||
for s in [
|
||||
server.alias,
|
||||
server.server_name,
|
||||
server.server_id,
|
||||
]
|
||||
if s is not None
|
||||
s.lower() for s in iter_known_server_prefixes(server) if s
|
||||
]
|
||||
|
||||
if server_or_group.lower() in match_list:
|
||||
|
|
@ -2031,11 +2026,13 @@ if MCP_AVAILABLE:
|
|||
# Remove prefix from tool name for logging and processing
|
||||
original_tool_name, server_name = split_server_prefix_from_name(name)
|
||||
|
||||
# If tool name is unprefixed, resolve its server so we can enforce permissions
|
||||
if not server_name:
|
||||
mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
|
||||
if mcp_server:
|
||||
server_name = mcp_server.name
|
||||
# Resolve the actual MCP server up-front so the permission check uses
|
||||
# the canonical server.name even when the tool name is prefixed with a
|
||||
# short ID (LITELLM_USE_SHORT_MCP_TOOL_PREFIX) that doesn't match the
|
||||
# server's display name directly.
|
||||
mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
|
||||
if mcp_server is not None:
|
||||
server_name = mcp_server.name
|
||||
|
||||
# Only enforce server-level permissions when we can resolve a server
|
||||
if server_name:
|
||||
|
|
|
|||
|
|
@ -2,10 +2,11 @@
|
|||
MCP Server Utilities
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, Mapping, Optional, Tuple
|
||||
from typing import Any, Dict, Iterator, Mapping, Optional, Tuple
|
||||
|
||||
import os
|
||||
import hashlib
|
||||
import importlib
|
||||
import os
|
||||
|
||||
# Constants
|
||||
LITELLM_MCP_SERVER_NAME = "litellm-mcp-server"
|
||||
|
|
@ -14,6 +15,89 @@ LITELLM_MCP_SERVER_DESCRIPTION = "MCP Server for LiteLLM"
|
|||
MCP_TOOL_PREFIX_SEPARATOR = os.environ.get("MCP_TOOL_PREFIX_SEPARATOR", "-")
|
||||
MCP_TOOL_PREFIX_FORMAT = "{server_name}{separator}{tool_name}"
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Short-ID tool prefix (opt-in)
|
||||
# ---------------------------------------------------------------------------
|
||||
# When LITELLM_USE_SHORT_MCP_TOOL_PREFIX is truthy the prefix attached to MCP
|
||||
# tool / prompt / resource / resource-template names switches from the
|
||||
# (potentially long) human-readable server name to a deterministic three
|
||||
# character ID derived from the server's ``server_id``.
|
||||
#
|
||||
# Why three characters?
|
||||
# * The first character is restricted to 52 alphabetic characters
|
||||
# ([A-Za-z]) and the remaining two characters use the full base62
|
||||
# alphabet ([0-9A-Za-z]). That guarantees the prefix never starts
|
||||
# with a digit so it remains a valid identifier for every model API
|
||||
# (some providers historically required a leading alphabetic char).
|
||||
# * 52 * 62 * 62 = 199_888 distinct IDs. The chance of a real local
|
||||
# tool name happening to begin with the exact prefix LiteLLM assigned
|
||||
# to a given MCP server is negligible in practice.
|
||||
# * The IDs are short enough that prefixed tool names stay well under
|
||||
# the 60-character upper bound enforced by some model APIs (Anthropic
|
||||
# etc.) even for long upstream tool names.
|
||||
# * The mapping is deterministic (SHA-256 of ``server_id`` → three
|
||||
# characters drawn from the alphabets above), so the prefix is stable
|
||||
# across processes, workers and restarts without any persistence
|
||||
# layer. Two servers with different ``server_id`` values can in
|
||||
# principle hash to the same three chars; that natural-hash collision
|
||||
# IS a routing-correctness issue (the second registrant would otherwise
|
||||
# have its tools misrouted to the first), so registration goes through
|
||||
# ``MCPServerManager._assign_unique_short_prefix`` which rehashes with
|
||||
# a deterministic attempt counter until it finds an unused prefix and
|
||||
# caches the result on ``MCPServer.short_prefix``. A collision is
|
||||
# logged at INFO when it happens.
|
||||
#
|
||||
# This flag is intentionally opt-in for the first release so customers can
|
||||
# migrate. It will become the default in a future release.
|
||||
SHORT_MCP_TOOL_PREFIX_LENGTH = 3
|
||||
_BASE62_ALPHABET = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz"
|
||||
# Subset of _BASE62_ALPHABET used for the *first* character only, to
|
||||
# guarantee the prefix never starts with a digit.
|
||||
_BASE52_ALPHA_ALPHABET = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz"
|
||||
|
||||
|
||||
def is_short_mcp_tool_prefix_enabled() -> bool:
|
||||
"""Return True when the short-ID tool prefix mode is enabled.
|
||||
|
||||
Read at call time (not import time) so tests and runtime config changes
|
||||
take effect without reimporting the module.
|
||||
"""
|
||||
raw = os.environ.get("LITELLM_USE_SHORT_MCP_TOOL_PREFIX", "")
|
||||
return raw.strip().lower() in ("1", "true", "yes", "on")
|
||||
|
||||
|
||||
def compute_short_server_prefix(server_id: str, attempt: int = 0) -> str:
|
||||
"""Derive the deterministic three-character prefix for a server.
|
||||
|
||||
Uses SHA-256 of ``f"{server_id}#{attempt}"`` and folds the first eight
|
||||
bytes into a fixed-length string whose first character is drawn from
|
||||
``_BASE52_ALPHA_ALPHABET`` (so the prefix never starts with a digit)
|
||||
and whose remaining characters are drawn from the full base62
|
||||
alphabet. Pass ``attempt > 0`` to rehash to a different prefix when
|
||||
the natural hash collides with a prefix already assigned to another
|
||||
server (see ``MCPServerManager._assign_unique_short_prefix``). An
|
||||
empty ``server_id`` raises ``ValueError`` — short prefixes require a
|
||||
stable identifier to be deterministic.
|
||||
"""
|
||||
if not server_id:
|
||||
raise ValueError("compute_short_server_prefix requires a non-empty server_id")
|
||||
|
||||
seed = server_id if attempt == 0 else f"{server_id}#{attempt}"
|
||||
digest = hashlib.sha256(seed.encode("utf-8")).digest()
|
||||
value = int.from_bytes(digest[:8], "big")
|
||||
|
||||
# Build chars from least-significant to most-significant; we reverse
|
||||
# at the end so the first emitted char comes from the high-order
|
||||
# bits of the digest (which is the position we constrain to be
|
||||
# alphabetic).
|
||||
chars = []
|
||||
for position in range(SHORT_MCP_TOOL_PREFIX_LENGTH):
|
||||
is_first_char = position == SHORT_MCP_TOOL_PREFIX_LENGTH - 1
|
||||
alphabet = _BASE52_ALPHA_ALPHABET if is_first_char else _BASE62_ALPHABET
|
||||
value, idx = divmod(value, len(alphabet))
|
||||
chars.append(alphabet[idx])
|
||||
return "".join(reversed(chars))
|
||||
|
||||
|
||||
def is_mcp_available() -> bool:
|
||||
"""
|
||||
|
|
@ -82,7 +166,25 @@ def add_server_prefix_to_name(name: str, server_name: str) -> str:
|
|||
|
||||
|
||||
def get_server_prefix(server: Any) -> str:
|
||||
"""Return the prefix for a server: alias if present, else server_name, else server_id"""
|
||||
"""Return the prefix for a server.
|
||||
|
||||
When the short-prefix mode is enabled (``LITELLM_USE_SHORT_MCP_TOOL_PREFIX``)
|
||||
a three-character base62 ID is returned. We prefer the cached
|
||||
``server.short_prefix`` value when set — that field is populated at
|
||||
registration time by ``MCPServerManager._assign_unique_short_prefix``
|
||||
and resolves natural-hash collisions deterministically — and only fall
|
||||
back to the natural hash for ad-hoc / temp-server objects without a
|
||||
cached value. In default mode the historical behaviour is preserved:
|
||||
alias if present, else server_name, else server_id.
|
||||
"""
|
||||
if is_short_mcp_tool_prefix_enabled():
|
||||
cached = getattr(server, "short_prefix", None)
|
||||
if cached:
|
||||
return cached
|
||||
server_id = getattr(server, "server_id", None)
|
||||
if server_id:
|
||||
return compute_short_server_prefix(server_id)
|
||||
|
||||
if hasattr(server, "alias") and server.alias:
|
||||
return server.alias
|
||||
if hasattr(server, "server_name") and server.server_name:
|
||||
|
|
@ -92,6 +194,36 @@ def get_server_prefix(server: Any) -> str:
|
|||
return ""
|
||||
|
||||
|
||||
def iter_known_server_prefixes(server: Any) -> Iterator[str]:
|
||||
"""Yield every prefix form that may appear in tool names for ``server``.
|
||||
|
||||
Always includes the *current* prefix returned by ``get_server_prefix``.
|
||||
Additionally yields the historical (alias / server_name / server_id) and
|
||||
short-ID forms so the routing layer can resolve tool names regardless of
|
||||
which prefix mode was active when the client first observed them.
|
||||
"""
|
||||
seen = set()
|
||||
|
||||
def _emit(value: Optional[str]) -> Iterator[str]:
|
||||
if value and value not in seen:
|
||||
seen.add(value)
|
||||
yield value
|
||||
|
||||
yield from _emit(get_server_prefix(server))
|
||||
yield from _emit(getattr(server, "short_prefix", None))
|
||||
|
||||
server_id = getattr(server, "server_id", None)
|
||||
if server_id:
|
||||
try:
|
||||
yield from _emit(compute_short_server_prefix(server_id))
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
yield from _emit(getattr(server, "alias", None))
|
||||
yield from _emit(getattr(server, "server_name", None))
|
||||
yield from _emit(server_id)
|
||||
|
||||
|
||||
def split_server_prefix_from_name(prefixed_name: str) -> Tuple[str, str]:
|
||||
"""Return the unprefixed name plus the server name used as prefix."""
|
||||
if MCP_TOOL_PREFIX_SEPARATOR in prefixed_name:
|
||||
|
|
|
|||
File diff suppressed because one or more lines are too long
|
|
@ -1,22 +0,0 @@
|
|||
1:"$Sreact.fragment"
|
||||
2:I[867271,["/litellm-asset-prefix/_next/static/chunks/9e09de50158b3159.js","/litellm-asset-prefix/_next/static/chunks/7e5fe5584502da06.js"],"default"]
|
||||
3:I[71195,["/litellm-asset-prefix/_next/static/chunks/9e09de50158b3159.js","/litellm-asset-prefix/_next/static/chunks/7e5fe5584502da06.js"],"default"]
|
||||
4:I[339756,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"default"]
|
||||
5:I[837457,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"default"]
|
||||
6:I[347257,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"ClientPageRoot"]
|
||||
7:I[321443,["/litellm-asset-prefix/_next/static/chunks/9e09de50158b3159.js","/litellm-asset-prefix/_next/static/chunks/7e5fe5584502da06.js","/litellm-asset-prefix/_next/static/chunks/0a6c418370a8c183.js","/litellm-asset-prefix/_next/static/chunks/cf9c81fc7166f4d4.js","/litellm-asset-prefix/_next/static/chunks/542a1a209eb732c6.js","/litellm-asset-prefix/_next/static/chunks/00ff280cdb7d7ee5.js","/litellm-asset-prefix/_next/static/chunks/a1abfc2f35c701cc.js","/litellm-asset-prefix/_next/static/chunks/4980372eaa37b78b.js","/litellm-asset-prefix/_next/static/chunks/119ecee91f911bc8.js","/litellm-asset-prefix/_next/static/chunks/ac3cf77acb5bf234.js","/litellm-asset-prefix/_next/static/chunks/fcdf7322b0aa3e2e.js","/litellm-asset-prefix/_next/static/chunks/7e417dd24c8becd0.js","/litellm-asset-prefix/_next/static/chunks/89aa55578de861b7.js"],"default"]
|
||||
a:I[897367,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"OutletBoundary"]
|
||||
b:"$Sreact.suspense"
|
||||
d:I[897367,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"ViewportBoundary"]
|
||||
f:I[897367,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"MetadataBoundary"]
|
||||
11:I[168027,[],"default"]
|
||||
:HL["/litellm-asset-prefix/_next/static/chunks/4e20891f2fd03463.css","style"]
|
||||
:HL["/litellm-asset-prefix/_next/static/chunks/e0cb6755699177c1.css","style"]
|
||||
:HL["/litellm-asset-prefix/_next/static/media/83afe278b6a6bb3c-s.p.3a6ba036.woff2","font",{"crossOrigin":"","type":"font/woff2"}]
|
||||
0:{"P":null,"b":"zxkD4-EPlgfKHDTw8O869","c":["","chat"],"q":"","i":false,"f":[[["",{"children":["chat",{"children":["__PAGE__",{}]}]},"$undefined","$undefined",true],[["$","$1","c",{"children":[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/chunks/4e20891f2fd03463.css","precedence":"next","crossOrigin":"$undefined","nonce":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/chunks/e0cb6755699177c1.css","precedence":"next","crossOrigin":"$undefined","nonce":"$undefined"}],["$","script","script-0",{"src":"/litellm-asset-prefix/_next/static/chunks/9e09de50158b3159.js","async":true,"nonce":"$undefined"}],["$","script","script-1",{"src":"/litellm-asset-prefix/_next/static/chunks/7e5fe5584502da06.js","async":true,"nonce":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"inter_5972bc34-module__OU16Qa__className","children":["$","$L2",null,{"children":["$","$L3",null,{"children":["$","$L4",null,{"parallelRouterKey":"children","error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":404}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],[]],"forbidden":"$undefined","unauthorized":"$undefined"}]}]}]}]}]]}],{"children":[["$","$1","c",{"children":[null,["$","$L4",null,{"parallelRouterKey":"children","error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":"$undefined","forbidden":"$undefined","unauthorized":"$undefined"}]]}],{"children":[["$","$1","c",{"children":[["$","$L6",null,{"Component":"$7","serverProvidedParams":{"searchParams":{},"params":{},"promises":["$@8","$@9"]}}],[["$","script","script-0",{"src":"/litellm-asset-prefix/_next/static/chunks/0a6c418370a8c183.js","async":true,"nonce":"$undefined"}],["$","script","script-1",{"src":"/litellm-asset-prefix/_next/static/chunks/cf9c81fc7166f4d4.js","async":true,"nonce":"$undefined"}],["$","script","script-2",{"src":"/litellm-asset-prefix/_next/static/chunks/542a1a209eb732c6.js","async":true,"nonce":"$undefined"}],["$","script","script-3",{"src":"/litellm-asset-prefix/_next/static/chunks/00ff280cdb7d7ee5.js","async":true,"nonce":"$undefined"}],["$","script","script-4",{"src":"/litellm-asset-prefix/_next/static/chunks/a1abfc2f35c701cc.js","async":true,"nonce":"$undefined"}],["$","script","script-5",{"src":"/litellm-asset-prefix/_next/static/chunks/4980372eaa37b78b.js","async":true,"nonce":"$undefined"}],["$","script","script-6",{"src":"/litellm-asset-prefix/_next/static/chunks/119ecee91f911bc8.js","async":true,"nonce":"$undefined"}],["$","script","script-7",{"src":"/litellm-asset-prefix/_next/static/chunks/ac3cf77acb5bf234.js","async":true,"nonce":"$undefined"}],["$","script","script-8",{"src":"/litellm-asset-prefix/_next/static/chunks/fcdf7322b0aa3e2e.js","async":true,"nonce":"$undefined"}],["$","script","script-9",{"src":"/litellm-asset-prefix/_next/static/chunks/7e417dd24c8becd0.js","async":true,"nonce":"$undefined"}],["$","script","script-10",{"src":"/litellm-asset-prefix/_next/static/chunks/89aa55578de861b7.js","async":true,"nonce":"$undefined"}]],["$","$La",null,{"children":["$","$b",null,{"name":"Next.MetadataOutlet","children":"$@c"}]}]]}],{},null,false,false]},null,false,false]},null,false,false],["$","$1","h",{"children":[null,["$","$Ld",null,{"children":"$Le"}],["$","div",null,{"hidden":true,"children":["$","$Lf",null,{"children":["$","$b",null,{"name":"Next.Metadata","children":"$L10"}]}]}],["$","meta",null,{"name":"next-size-adjust","content":""}]]}],false]],"m":"$undefined","G":["$11",[]],"S":true}
|
||||
8:{}
|
||||
9:"$0:f:0:1:1:children:1:children:0:props:children:0:props:serverProvidedParams:params"
|
||||
e:[["$","meta","0",{"charSet":"utf-8"}],["$","meta","1",{"name":"viewport","content":"width=device-width, initial-scale=1"}]]
|
||||
12:I[27201,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"IconMark"]
|
||||
c:null
|
||||
10:[["$","title","0",{"children":"LiteLLM Dashboard"}],["$","meta","1",{"name":"description","content":"LiteLLM Proxy Admin UI"}],["$","link","2",{"rel":"icon","href":"/favicon.ico?favicon.1d32c690.ico","sizes":"48x48","type":"image/x-icon"}],["$","link","3",{"rel":"icon","href":"./favicon.ico"}],["$","$L12","4",{}]]
|
||||
|
|
@ -1,22 +0,0 @@
|
|||
1:"$Sreact.fragment"
|
||||
2:I[867271,["/litellm-asset-prefix/_next/static/chunks/9e09de50158b3159.js","/litellm-asset-prefix/_next/static/chunks/7e5fe5584502da06.js"],"default"]
|
||||
3:I[71195,["/litellm-asset-prefix/_next/static/chunks/9e09de50158b3159.js","/litellm-asset-prefix/_next/static/chunks/7e5fe5584502da06.js"],"default"]
|
||||
4:I[339756,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"default"]
|
||||
5:I[837457,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"default"]
|
||||
6:I[347257,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"ClientPageRoot"]
|
||||
7:I[321443,["/litellm-asset-prefix/_next/static/chunks/9e09de50158b3159.js","/litellm-asset-prefix/_next/static/chunks/7e5fe5584502da06.js","/litellm-asset-prefix/_next/static/chunks/0a6c418370a8c183.js","/litellm-asset-prefix/_next/static/chunks/cf9c81fc7166f4d4.js","/litellm-asset-prefix/_next/static/chunks/542a1a209eb732c6.js","/litellm-asset-prefix/_next/static/chunks/00ff280cdb7d7ee5.js","/litellm-asset-prefix/_next/static/chunks/a1abfc2f35c701cc.js","/litellm-asset-prefix/_next/static/chunks/4980372eaa37b78b.js","/litellm-asset-prefix/_next/static/chunks/119ecee91f911bc8.js","/litellm-asset-prefix/_next/static/chunks/ac3cf77acb5bf234.js","/litellm-asset-prefix/_next/static/chunks/fcdf7322b0aa3e2e.js","/litellm-asset-prefix/_next/static/chunks/7e417dd24c8becd0.js","/litellm-asset-prefix/_next/static/chunks/89aa55578de861b7.js"],"default"]
|
||||
a:I[897367,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"OutletBoundary"]
|
||||
b:"$Sreact.suspense"
|
||||
d:I[897367,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"ViewportBoundary"]
|
||||
f:I[897367,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"MetadataBoundary"]
|
||||
11:I[168027,[],"default"]
|
||||
:HL["/litellm-asset-prefix/_next/static/chunks/4e20891f2fd03463.css","style"]
|
||||
:HL["/litellm-asset-prefix/_next/static/chunks/e0cb6755699177c1.css","style"]
|
||||
:HL["/litellm-asset-prefix/_next/static/media/83afe278b6a6bb3c-s.p.3a6ba036.woff2","font",{"crossOrigin":"","type":"font/woff2"}]
|
||||
0:{"P":null,"b":"zxkD4-EPlgfKHDTw8O869","c":["","chat"],"q":"","i":false,"f":[[["",{"children":["chat",{"children":["__PAGE__",{}]}]},"$undefined","$undefined",true],[["$","$1","c",{"children":[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/chunks/4e20891f2fd03463.css","precedence":"next","crossOrigin":"$undefined","nonce":"$undefined"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/chunks/e0cb6755699177c1.css","precedence":"next","crossOrigin":"$undefined","nonce":"$undefined"}],["$","script","script-0",{"src":"/litellm-asset-prefix/_next/static/chunks/9e09de50158b3159.js","async":true,"nonce":"$undefined"}],["$","script","script-1",{"src":"/litellm-asset-prefix/_next/static/chunks/7e5fe5584502da06.js","async":true,"nonce":"$undefined"}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"inter_5972bc34-module__OU16Qa__className","children":["$","$L2",null,{"children":["$","$L3",null,{"children":["$","$L4",null,{"parallelRouterKey":"children","error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":[[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":404}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],[]],"forbidden":"$undefined","unauthorized":"$undefined"}]}]}]}]}]]}],{"children":[["$","$1","c",{"children":[null,["$","$L4",null,{"parallelRouterKey":"children","error":"$undefined","errorStyles":"$undefined","errorScripts":"$undefined","template":["$","$L5",null,{}],"templateStyles":"$undefined","templateScripts":"$undefined","notFound":"$undefined","forbidden":"$undefined","unauthorized":"$undefined"}]]}],{"children":[["$","$1","c",{"children":[["$","$L6",null,{"Component":"$7","serverProvidedParams":{"searchParams":{},"params":{},"promises":["$@8","$@9"]}}],[["$","script","script-0",{"src":"/litellm-asset-prefix/_next/static/chunks/0a6c418370a8c183.js","async":true,"nonce":"$undefined"}],["$","script","script-1",{"src":"/litellm-asset-prefix/_next/static/chunks/cf9c81fc7166f4d4.js","async":true,"nonce":"$undefined"}],["$","script","script-2",{"src":"/litellm-asset-prefix/_next/static/chunks/542a1a209eb732c6.js","async":true,"nonce":"$undefined"}],["$","script","script-3",{"src":"/litellm-asset-prefix/_next/static/chunks/00ff280cdb7d7ee5.js","async":true,"nonce":"$undefined"}],["$","script","script-4",{"src":"/litellm-asset-prefix/_next/static/chunks/a1abfc2f35c701cc.js","async":true,"nonce":"$undefined"}],["$","script","script-5",{"src":"/litellm-asset-prefix/_next/static/chunks/4980372eaa37b78b.js","async":true,"nonce":"$undefined"}],["$","script","script-6",{"src":"/litellm-asset-prefix/_next/static/chunks/119ecee91f911bc8.js","async":true,"nonce":"$undefined"}],["$","script","script-7",{"src":"/litellm-asset-prefix/_next/static/chunks/ac3cf77acb5bf234.js","async":true,"nonce":"$undefined"}],["$","script","script-8",{"src":"/litellm-asset-prefix/_next/static/chunks/fcdf7322b0aa3e2e.js","async":true,"nonce":"$undefined"}],["$","script","script-9",{"src":"/litellm-asset-prefix/_next/static/chunks/7e417dd24c8becd0.js","async":true,"nonce":"$undefined"}],["$","script","script-10",{"src":"/litellm-asset-prefix/_next/static/chunks/89aa55578de861b7.js","async":true,"nonce":"$undefined"}]],["$","$La",null,{"children":["$","$b",null,{"name":"Next.MetadataOutlet","children":"$@c"}]}]]}],{},null,false,false]},null,false,false]},null,false,false],["$","$1","h",{"children":[null,["$","$Ld",null,{"children":"$Le"}],["$","div",null,{"hidden":true,"children":["$","$Lf",null,{"children":["$","$b",null,{"name":"Next.Metadata","children":"$L10"}]}]}],["$","meta",null,{"name":"next-size-adjust","content":""}]]}],false]],"m":"$undefined","G":["$11",[]],"S":true}
|
||||
8:{}
|
||||
9:"$0:f:0:1:1:children:1:children:0:props:children:0:props:serverProvidedParams:params"
|
||||
e:[["$","meta","0",{"charSet":"utf-8"}],["$","meta","1",{"name":"viewport","content":"width=device-width, initial-scale=1"}]]
|
||||
12:I[27201,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"IconMark"]
|
||||
c:null
|
||||
10:[["$","title","0",{"children":"LiteLLM Dashboard"}],["$","meta","1",{"name":"description","content":"LiteLLM Proxy Admin UI"}],["$","link","2",{"rel":"icon","href":"/favicon.ico?favicon.1d32c690.ico","sizes":"48x48","type":"image/x-icon"}],["$","link","3",{"rel":"icon","href":"./favicon.ico"}],["$","$L12","4",{}]]
|
||||
|
|
@ -1,6 +0,0 @@
|
|||
1:"$Sreact.fragment"
|
||||
2:I[897367,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"ViewportBoundary"]
|
||||
3:I[897367,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"MetadataBoundary"]
|
||||
4:"$Sreact.suspense"
|
||||
5:I[27201,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"IconMark"]
|
||||
0:{"buildId":"zxkD4-EPlgfKHDTw8O869","rsc":["$","$1","h",{"children":[null,["$","$L2",null,{"children":[["$","meta","0",{"charSet":"utf-8"}],["$","meta","1",{"name":"viewport","content":"width=device-width, initial-scale=1"}]]}],["$","div",null,{"hidden":true,"children":["$","$L3",null,{"children":["$","$4",null,{"name":"Next.Metadata","children":[["$","title","0",{"children":"LiteLLM Dashboard"}],["$","meta","1",{"name":"description","content":"LiteLLM Proxy Admin UI"}],["$","link","2",{"rel":"icon","href":"/favicon.ico?favicon.1d32c690.ico","sizes":"48x48","type":"image/x-icon"}],["$","link","3",{"rel":"icon","href":"./favicon.ico"}],["$","$L5","4",{}]]}]}]}],["$","meta",null,{"name":"next-size-adjust","content":""}]]}],"loading":null,"isPartial":false}
|
||||
|
|
@ -1,8 +0,0 @@
|
|||
1:"$Sreact.fragment"
|
||||
2:I[867271,["/litellm-asset-prefix/_next/static/chunks/9e09de50158b3159.js","/litellm-asset-prefix/_next/static/chunks/7e5fe5584502da06.js"],"default"]
|
||||
3:I[71195,["/litellm-asset-prefix/_next/static/chunks/9e09de50158b3159.js","/litellm-asset-prefix/_next/static/chunks/7e5fe5584502da06.js"],"default"]
|
||||
4:I[339756,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"default"]
|
||||
5:I[837457,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"default"]
|
||||
:HL["/litellm-asset-prefix/_next/static/chunks/4e20891f2fd03463.css","style"]
|
||||
:HL["/litellm-asset-prefix/_next/static/chunks/e0cb6755699177c1.css","style"]
|
||||
0:{"buildId":"zxkD4-EPlgfKHDTw8O869","rsc":["$","$1","c",{"children":[[["$","link","0",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/chunks/4e20891f2fd03463.css","precedence":"next"}],["$","link","1",{"rel":"stylesheet","href":"/litellm-asset-prefix/_next/static/chunks/e0cb6755699177c1.css","precedence":"next"}],["$","script","script-0",{"src":"/litellm-asset-prefix/_next/static/chunks/9e09de50158b3159.js","async":true}],["$","script","script-1",{"src":"/litellm-asset-prefix/_next/static/chunks/7e5fe5584502da06.js","async":true}]],["$","html",null,{"lang":"en","children":["$","body",null,{"className":"inter_5972bc34-module__OU16Qa__className","children":["$","$L2",null,{"children":["$","$L3",null,{"children":["$","$L4",null,{"parallelRouterKey":"children","template":["$","$L5",null,{}],"notFound":[[["$","title",null,{"children":"404: This page could not be found."}],["$","div",null,{"style":{"fontFamily":"system-ui,\"Segoe UI\",Roboto,Helvetica,Arial,sans-serif,\"Apple Color Emoji\",\"Segoe UI Emoji\"","height":"100vh","textAlign":"center","display":"flex","flexDirection":"column","alignItems":"center","justifyContent":"center"},"children":["$","div",null,{"children":[["$","style",null,{"dangerouslySetInnerHTML":{"__html":"body{color:#000;background:#fff;margin:0}.next-error-h1{border-right:1px solid rgba(0,0,0,.3)}@media (prefers-color-scheme:dark){body{color:#fff;background:#000}.next-error-h1{border-right:1px solid rgba(255,255,255,.3)}}"}}],["$","h1",null,{"className":"next-error-h1","style":{"display":"inline-block","margin":"0 20px 0 0","padding":"0 23px 0 0","fontSize":24,"fontWeight":500,"verticalAlign":"top","lineHeight":"49px"},"children":404}],["$","div",null,{"style":{"display":"inline-block"},"children":["$","h2",null,{"style":{"fontSize":14,"fontWeight":400,"lineHeight":"49px","margin":0},"children":"This page could not be found."}]}]]}]}]],[]]}]}]}]}]}]]}],"loading":null,"isPartial":false}
|
||||
|
|
@ -1,4 +0,0 @@
|
|||
:HL["/litellm-asset-prefix/_next/static/chunks/4e20891f2fd03463.css","style"]
|
||||
:HL["/litellm-asset-prefix/_next/static/chunks/e0cb6755699177c1.css","style"]
|
||||
:HL["/litellm-asset-prefix/_next/static/media/83afe278b6a6bb3c-s.p.3a6ba036.woff2","font",{"crossOrigin":"","type":"font/woff2"}]
|
||||
0:{"buildId":"zxkD4-EPlgfKHDTw8O869","tree":{"name":"","paramType":null,"paramKey":"","hasRuntimePrefetch":false,"slots":{"children":{"name":"chat","paramType":null,"paramKey":"chat","hasRuntimePrefetch":false,"slots":{"children":{"name":"__PAGE__","paramType":null,"paramKey":"__PAGE__","hasRuntimePrefetch":false,"slots":null,"isRootLayout":false}},"isRootLayout":false}},"isRootLayout":true},"staleTime":300}
|
||||
|
|
@ -1,9 +0,0 @@
|
|||
1:"$Sreact.fragment"
|
||||
2:I[347257,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"ClientPageRoot"]
|
||||
3:I[321443,["/litellm-asset-prefix/_next/static/chunks/9e09de50158b3159.js","/litellm-asset-prefix/_next/static/chunks/7e5fe5584502da06.js","/litellm-asset-prefix/_next/static/chunks/0a6c418370a8c183.js","/litellm-asset-prefix/_next/static/chunks/cf9c81fc7166f4d4.js","/litellm-asset-prefix/_next/static/chunks/542a1a209eb732c6.js","/litellm-asset-prefix/_next/static/chunks/00ff280cdb7d7ee5.js","/litellm-asset-prefix/_next/static/chunks/a1abfc2f35c701cc.js","/litellm-asset-prefix/_next/static/chunks/4980372eaa37b78b.js","/litellm-asset-prefix/_next/static/chunks/119ecee91f911bc8.js","/litellm-asset-prefix/_next/static/chunks/ac3cf77acb5bf234.js","/litellm-asset-prefix/_next/static/chunks/fcdf7322b0aa3e2e.js","/litellm-asset-prefix/_next/static/chunks/7e417dd24c8becd0.js","/litellm-asset-prefix/_next/static/chunks/89aa55578de861b7.js"],"default"]
|
||||
6:I[897367,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"OutletBoundary"]
|
||||
7:"$Sreact.suspense"
|
||||
0:{"buildId":"zxkD4-EPlgfKHDTw8O869","rsc":["$","$1","c",{"children":[["$","$L2",null,{"Component":"$3","serverProvidedParams":{"searchParams":{},"params":{},"promises":["$@4","$@5"]}}],[["$","script","script-0",{"src":"/litellm-asset-prefix/_next/static/chunks/0a6c418370a8c183.js","async":true}],["$","script","script-1",{"src":"/litellm-asset-prefix/_next/static/chunks/cf9c81fc7166f4d4.js","async":true}],["$","script","script-2",{"src":"/litellm-asset-prefix/_next/static/chunks/542a1a209eb732c6.js","async":true}],["$","script","script-3",{"src":"/litellm-asset-prefix/_next/static/chunks/00ff280cdb7d7ee5.js","async":true}],["$","script","script-4",{"src":"/litellm-asset-prefix/_next/static/chunks/a1abfc2f35c701cc.js","async":true}],["$","script","script-5",{"src":"/litellm-asset-prefix/_next/static/chunks/4980372eaa37b78b.js","async":true}],["$","script","script-6",{"src":"/litellm-asset-prefix/_next/static/chunks/119ecee91f911bc8.js","async":true}],["$","script","script-7",{"src":"/litellm-asset-prefix/_next/static/chunks/ac3cf77acb5bf234.js","async":true}],["$","script","script-8",{"src":"/litellm-asset-prefix/_next/static/chunks/fcdf7322b0aa3e2e.js","async":true}],["$","script","script-9",{"src":"/litellm-asset-prefix/_next/static/chunks/7e417dd24c8becd0.js","async":true}],["$","script","script-10",{"src":"/litellm-asset-prefix/_next/static/chunks/89aa55578de861b7.js","async":true}]],["$","$L6",null,{"children":["$","$7",null,{"name":"Next.MetadataOutlet","children":"$@8"}]}]]}],"loading":null,"isPartial":false}
|
||||
4:{}
|
||||
5:"$0:rsc:props:children:0:props:serverProvidedParams:params"
|
||||
8:null
|
||||
|
|
@ -1,4 +0,0 @@
|
|||
1:"$Sreact.fragment"
|
||||
2:I[339756,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"default"]
|
||||
3:I[837457,["/litellm-asset-prefix/_next/static/chunks/d96012bcfc98706a.js","/litellm-asset-prefix/_next/static/chunks/dbca964212122d58.js"],"default"]
|
||||
0:{"buildId":"zxkD4-EPlgfKHDTw8O869","rsc":["$","$1","c",{"children":[null,["$","$L2",null,{"parallelRouterKey":"children","template":["$","$L3",null,{}]}]]}],"loading":null,"isPartial":false}
|
||||
|
|
@ -619,6 +619,67 @@ class ProxyBaseLLMRequestProcessing:
|
|||
verbose_proxy_logger.error(f"Error setting custom headers: {e}")
|
||||
return {}
|
||||
|
||||
@staticmethod
|
||||
async def build_litellm_proxy_success_headers_from_llm_response(
|
||||
*,
|
||||
response: Any,
|
||||
request_data: dict,
|
||||
request: Request,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
version: Optional[str],
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
) -> Dict[str, str]:
|
||||
"""
|
||||
Build LiteLLM proxy response headers for routes that call the LLM directly
|
||||
(e.g. Google native :generateContent) instead of base_process_llm_request.
|
||||
"""
|
||||
if isinstance(response, dict):
|
||||
hidden_params = response.get("_hidden_params") or {}
|
||||
else:
|
||||
hidden_params = getattr(response, "_hidden_params", None) or {}
|
||||
if not isinstance(hidden_params, dict):
|
||||
hidden_params = {}
|
||||
|
||||
model_id = ProxyBaseLLMRequestProcessing._get_model_id_from_response(
|
||||
hidden_params, request_data
|
||||
)
|
||||
|
||||
cache_key = hidden_params.get("cache_key", None) or ""
|
||||
api_base = hidden_params.get("api_base", None) or ""
|
||||
response_cost = hidden_params.get("response_cost", None) or ""
|
||||
fastest_response_batch_completion = hidden_params.get(
|
||||
"fastest_response_batch_completion", None
|
||||
)
|
||||
additional_headers = hidden_params.get("additional_headers", {}) or {}
|
||||
|
||||
custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_id=logging_obj.litellm_call_id,
|
||||
model_id=model_id,
|
||||
cache_key=cache_key,
|
||||
api_base=api_base,
|
||||
version=version,
|
||||
response_cost=response_cost,
|
||||
model_region=getattr(user_api_key_dict, "allowed_model_region", ""),
|
||||
fastest_response_batch_completion=fastest_response_batch_completion,
|
||||
request_data=request_data,
|
||||
hidden_params=hidden_params,
|
||||
litellm_logging_obj=logging_obj,
|
||||
**additional_headers,
|
||||
)
|
||||
|
||||
callback_headers = await proxy_logging_obj.post_call_response_headers_hook(
|
||||
data=request_data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=response,
|
||||
request_headers=dict(request.headers),
|
||||
)
|
||||
if callback_headers:
|
||||
custom_headers.update(callback_headers)
|
||||
|
||||
return custom_headers
|
||||
|
||||
async def common_processing_pre_call_logic(
|
||||
self,
|
||||
request: Request,
|
||||
|
|
@ -875,7 +936,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
else:
|
||||
verbose_proxy_logger.debug(
|
||||
"Request received by LiteLLM:\n%s",
|
||||
json.dumps(self.data, indent=4, default=str),
|
||||
_payload_str,
|
||||
)
|
||||
|
||||
async def base_process_llm_request( # noqa: PLR0915
|
||||
|
|
@ -1511,9 +1572,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
_response = assembled_response
|
||||
try:
|
||||
from litellm.proxy.proxy_server import llm_router as _global_llm_router
|
||||
from litellm.proxy.utils import (
|
||||
_check_and_merge_model_level_guardrails,
|
||||
)
|
||||
from litellm.proxy.utils import _check_and_merge_model_level_guardrails
|
||||
|
||||
guardrail_data = _check_and_merge_model_level_guardrails(
|
||||
data=captured_data, llm_router=_global_llm_router
|
||||
|
|
@ -1690,11 +1749,12 @@ class ProxyBaseLLMRequestProcessing:
|
|||
elif isinstance(e, httpx.HTTPStatusError):
|
||||
# Handle httpx.HTTPStatusError - extract actual error from response
|
||||
# This matches the original behavior before the refactor in commit 511d435f6f
|
||||
error_body = await e.response.aread()
|
||||
http_status_error: httpx.HTTPStatusError = e
|
||||
error_body = await http_status_error.response.aread()
|
||||
error_text = error_body.decode("utf-8")
|
||||
|
||||
raise HTTPException(
|
||||
status_code=e.response.status_code,
|
||||
status_code=http_status_error.response.status_code,
|
||||
detail={"error": error_text},
|
||||
)
|
||||
error_msg = f"{str(e)}"
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
from typing import Union
|
||||
from typing import Any, Awaitable, Callable, Optional, Union
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import (
|
||||
DB_CONNECTION_ERROR_TYPES,
|
||||
ProxyErrorTypes,
|
||||
|
|
@ -123,3 +124,138 @@ class PrismaDBExceptionHandler:
|
|||
):
|
||||
return None
|
||||
raise e
|
||||
|
||||
|
||||
# Default fallback timeouts when neither the caller nor the prisma_client
|
||||
# expose `_db_auth_reconnect_timeout_seconds` / `_db_auth_reconnect_lock_timeout_seconds`.
|
||||
# Match the auth path's existing defaults so behavior is uniform across read paths.
|
||||
_DEFAULT_RECONNECT_TIMEOUT_SECONDS = 2.0
|
||||
_DEFAULT_RECONNECT_LOCK_TIMEOUT_SECONDS = 0.1
|
||||
|
||||
|
||||
def _coerce_timeout(value: Any, fallback: float) -> float:
|
||||
"""Return `value` if it is a real int/float, else `fallback`. Guards
|
||||
against tests that mock `prisma_client` and leave the timeout slots as
|
||||
MagicMock instances."""
|
||||
if isinstance(value, (int, float)) and not isinstance(value, bool):
|
||||
return float(value)
|
||||
return fallback
|
||||
|
||||
|
||||
async def call_with_db_reconnect_retry(
|
||||
prisma_client: Any,
|
||||
coro_factory: Callable[[], Awaitable[Any]],
|
||||
*,
|
||||
reason: str,
|
||||
timeout_seconds: Optional[float] = None,
|
||||
lock_timeout_seconds: Optional[float] = None,
|
||||
) -> Any:
|
||||
"""Run a Prisma read coroutine with one transport-reconnect-and-retry.
|
||||
|
||||
The canonical "self-heal a transient DB transport blip" wrapper used by
|
||||
`PrismaClient.get_generic_data` and other read paths. Mirrors the inline
|
||||
pattern in `auth_checks._fetch_key_object_from_db_with_reconnect` so we
|
||||
have a single implementation rather than three drifting copies.
|
||||
|
||||
Behavior:
|
||||
1. Await `coro_factory()`. On success, return its value.
|
||||
2. On exception, if it is NOT a transport error (per
|
||||
`is_database_transport_error`), re-raise — data-layer errors like
|
||||
`UniqueViolationError` mean the DB is reachable, reconnect would be
|
||||
pointless.
|
||||
3. If `prisma_client` does not expose `attempt_db_reconnect`, re-raise.
|
||||
This guards against partial stand-ins / older clients in tests.
|
||||
4. Call `prisma_client.attempt_db_reconnect(reason=...)`. If it returns
|
||||
False (cooldown / lock contention / reconnect failure), re-raise.
|
||||
5. Otherwise await `coro_factory()` a second time and return / propagate
|
||||
its result. At-most-one retry by construction — no infinite loop.
|
||||
|
||||
`coro_factory` MUST be a zero-arg callable that returns a fresh awaitable
|
||||
on each call. Passing an already-awaited coroutine would fail on retry
|
||||
with `RuntimeError: cannot reuse already awaited coroutine`.
|
||||
|
||||
`reason` should follow `<subsystem>_<operation>_<table>_failure` so
|
||||
telemetry distinguishes between fan-out callers (e.g.
|
||||
`_update_config_from_db` issues four concurrent reads).
|
||||
|
||||
Args:
|
||||
prisma_client: The `PrismaClient` (or stand-in) that owns
|
||||
`attempt_db_reconnect` and the `_db_auth_reconnect_*` defaults.
|
||||
coro_factory: Zero-arg callable returning the read awaitable.
|
||||
reason: Telemetry tag forwarded to `attempt_db_reconnect`.
|
||||
timeout_seconds: Optional override for the reconnect cycle timeout.
|
||||
Defaults to `prisma_client._db_auth_reconnect_timeout_seconds`,
|
||||
then to 2.0s.
|
||||
lock_timeout_seconds: Optional override for how long the helper will
|
||||
wait to acquire the reconnect lock. Defaults to
|
||||
`prisma_client._db_auth_reconnect_lock_timeout_seconds`, then to
|
||||
0.1s.
|
||||
|
||||
Returns:
|
||||
Whatever `coro_factory()` returns (on first or second attempt).
|
||||
|
||||
Raises:
|
||||
Whatever `coro_factory()` raises if the failure is not a transport
|
||||
error, or if the reconnect attempt does not succeed, or if the retry
|
||||
also fails.
|
||||
"""
|
||||
try:
|
||||
return await coro_factory()
|
||||
except Exception as first_exc:
|
||||
if not PrismaDBExceptionHandler.is_database_transport_error(first_exc):
|
||||
raise
|
||||
if not hasattr(prisma_client, "attempt_db_reconnect"):
|
||||
raise
|
||||
|
||||
resolved_timeout = _coerce_timeout(
|
||||
(
|
||||
timeout_seconds
|
||||
if timeout_seconds is not None
|
||||
else getattr(prisma_client, "_db_auth_reconnect_timeout_seconds", None)
|
||||
),
|
||||
_DEFAULT_RECONNECT_TIMEOUT_SECONDS,
|
||||
)
|
||||
resolved_lock_timeout = _coerce_timeout(
|
||||
(
|
||||
lock_timeout_seconds
|
||||
if lock_timeout_seconds is not None
|
||||
else getattr(
|
||||
prisma_client, "_db_auth_reconnect_lock_timeout_seconds", None
|
||||
)
|
||||
),
|
||||
_DEFAULT_RECONNECT_LOCK_TIMEOUT_SECONDS,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.warning(
|
||||
"DB transport error on read; attempting reconnect-and-retry. reason=%s error=%s",
|
||||
reason,
|
||||
first_exc,
|
||||
)
|
||||
|
||||
# Preserve the original transport error in telemetry. If
|
||||
# `attempt_db_reconnect` itself raises (e.g. lock cancellation, timer
|
||||
# error, unexpected internal failure), surfacing that exception
|
||||
# instead of `first_exc` would mask the actual DB transport problem
|
||||
# in `failure_handler` / `db_exceptions` alerts. Chain the reconnect
|
||||
# error as the cause for debuggability without losing the original.
|
||||
try:
|
||||
did_reconnect = await prisma_client.attempt_db_reconnect(
|
||||
reason=reason,
|
||||
timeout_seconds=resolved_timeout,
|
||||
lock_timeout_seconds=resolved_lock_timeout,
|
||||
)
|
||||
except Exception as reconnect_exc:
|
||||
verbose_proxy_logger.warning(
|
||||
"DB reconnect attempt raised; preserving original transport error. "
|
||||
"reason=%s reconnect_error=%s",
|
||||
reason,
|
||||
reconnect_exc,
|
||||
)
|
||||
raise first_exc from reconnect_exc
|
||||
if not did_reconnect:
|
||||
raise
|
||||
|
||||
# At most one retry. If the retry also raises a transport error, we
|
||||
# propagate — repeated reconnect-loops are the watchdog's job, not
|
||||
# this helper's.
|
||||
return await coro_factory()
|
||||
|
|
|
|||
|
|
@ -52,18 +52,25 @@ class PrismaWrapper:
|
|||
engine = self._original_prisma._engine
|
||||
process = getattr(engine, "process", None) if engine is not None else None
|
||||
if process is not None:
|
||||
return process.pid
|
||||
pid = process.pid
|
||||
if isinstance(pid, int):
|
||||
return pid
|
||||
except (AttributeError, TypeError):
|
||||
pass
|
||||
return 0
|
||||
|
||||
@staticmethod
|
||||
async def _kill_engine_process(pid: int) -> None:
|
||||
"""Force-kill an orphaned engine subprocess to prevent DB connection pool leaks.
|
||||
"""Force-kill the engine subprocess to prevent DB connection pool leaks.
|
||||
|
||||
Called when disconnect() fails and the old engine process may still be
|
||||
holding open connections. Sends SIGTERM for graceful shutdown, waits
|
||||
briefly, then SIGKILL as a backstop.
|
||||
Called on every reconnect (in `recreate_prisma_client`) to retire the
|
||||
old query-engine subprocess without invoking prisma-client-py's
|
||||
synchronous `disconnect()` — which blocks the asyncio event loop on
|
||||
`subprocess.Popen.wait()` for 30-120+ seconds when the engine is
|
||||
stuck on TCP close.
|
||||
|
||||
Sends SIGTERM for graceful shutdown, waits briefly, then SIGKILL as
|
||||
a backstop.
|
||||
"""
|
||||
if pid <= 0:
|
||||
return
|
||||
|
|
@ -72,7 +79,7 @@ class PrismaWrapper:
|
|||
except (ProcessLookupError, PermissionError, OSError):
|
||||
return # Already dead or inaccessible
|
||||
verbose_proxy_logger.warning(
|
||||
"Sent SIGTERM to orphaned prisma-query-engine PID %s after failed disconnect.",
|
||||
"Sent SIGTERM to prisma-query-engine PID %s during reconnect.",
|
||||
pid,
|
||||
)
|
||||
# Brief wait for graceful shutdown, then force-kill
|
||||
|
|
@ -217,15 +224,18 @@ class PrismaWrapper:
|
|||
async def recreate_prisma_client(
|
||||
self, new_db_url: str, http_client: Optional[Any] = None
|
||||
):
|
||||
"""Disconnect and reconnect the Prisma client with a new database URL."""
|
||||
"""Disconnect and reconnect the Prisma client with a new database URL.
|
||||
|
||||
Kills the old engine subprocess directly (SIGTERM → SIGKILL) rather than
|
||||
calling `disconnect()`. prisma-client-py's `disconnect()` calls a
|
||||
synchronous `subprocess.Popen.wait()` that can freeze the asyncio event
|
||||
loop for 30-120+ seconds when the engine is stuck on TCP close,
|
||||
breaking `/health/liveliness` and causing Kubernetes pod restarts.
|
||||
"""
|
||||
from prisma import Prisma # type: ignore
|
||||
|
||||
old_engine_pid = self._get_engine_pid()
|
||||
|
||||
try:
|
||||
await self._original_prisma.disconnect()
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(f"Failed to disconnect Prisma client: {e}")
|
||||
if old_engine_pid > 0:
|
||||
await self._kill_engine_process(old_engine_pid)
|
||||
|
||||
if http_client is not None:
|
||||
|
|
|
|||
|
|
@ -35,6 +35,7 @@ async def google_generate_content(
|
|||
general_settings,
|
||||
llm_router,
|
||||
proxy_config,
|
||||
proxy_logging_obj,
|
||||
version,
|
||||
)
|
||||
|
||||
|
|
@ -73,6 +74,16 @@ async def google_generate_content(
|
|||
if llm_router is None:
|
||||
raise HTTPException(status_code=500, detail="Router not initialized")
|
||||
response = await llm_router.agenerate_content(**data)
|
||||
success_headers = await ProxyBaseLLMRequestProcessing.build_litellm_proxy_success_headers_from_llm_response(
|
||||
response=response,
|
||||
request_data=data,
|
||||
request=request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
logging_obj=logging_obj,
|
||||
version=version,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
fastapi_response.headers.update(success_headers)
|
||||
return response
|
||||
|
||||
|
||||
|
|
@ -95,6 +106,7 @@ async def google_stream_generate_content(
|
|||
general_settings,
|
||||
llm_router,
|
||||
proxy_config,
|
||||
proxy_logging_obj,
|
||||
version,
|
||||
)
|
||||
|
||||
|
|
@ -137,9 +149,24 @@ async def google_stream_generate_content(
|
|||
raise HTTPException(status_code=500, detail="Router not initialized")
|
||||
response = await llm_router.agenerate_content_stream(**data)
|
||||
|
||||
success_headers = await ProxyBaseLLMRequestProcessing.build_litellm_proxy_success_headers_from_llm_response(
|
||||
response=response,
|
||||
request_data=data,
|
||||
request=request,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
logging_obj=logging_obj,
|
||||
version=version,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
# Check if response is an async iterator (streaming response)
|
||||
if response is not None and hasattr(response, "__aiter__"):
|
||||
return StreamingResponse(content=response, media_type="text/event-stream")
|
||||
return StreamingResponse(
|
||||
content=response,
|
||||
media_type="text/event-stream",
|
||||
headers=success_headers,
|
||||
)
|
||||
fastapi_response.headers.update(success_headers)
|
||||
return response
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,492 @@
|
|||
"""
|
||||
WORKFLOW RUN MANAGEMENT
|
||||
|
||||
Generic durable state tracking for agents and automated workflows.
|
||||
|
||||
POST /v1/workflows/runs - Create a workflow run
|
||||
GET /v1/workflows/runs - List runs (filter by type, status)
|
||||
GET /v1/workflows/runs/{run_id} - Get run with latest event
|
||||
PATCH /v1/workflows/runs/{run_id} - Update status, metadata, output
|
||||
POST /v1/workflows/runs/{run_id}/events - Append event (updates run status)
|
||||
GET /v1/workflows/runs/{run_id}/events - Full event log
|
||||
POST /v1/workflows/runs/{run_id}/messages - Append conversation message
|
||||
GET /v1/workflows/runs/{run_id}/messages - Fetch conversation history
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, Dict, Literal, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
|
||||
try:
|
||||
from prisma.errors import UniqueViolationError
|
||||
except ImportError:
|
||||
UniqueViolationError = None # type: ignore
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
_MAX_SEQUENCE_RETRIES = 5
|
||||
|
||||
|
||||
def _json(value: Any) -> str:
|
||||
"""Serialize a Python value for prisma-client-py Json fields (must be a string)."""
|
||||
return json.dumps(value)
|
||||
|
||||
|
||||
def _is_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
|
||||
return user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
|
||||
|
||||
|
||||
def _caller_key(user_api_key_dict: UserAPIKeyAuth) -> Optional[str]:
|
||||
"""Return the hashed key token that identifies this caller, or None for master key."""
|
||||
return user_api_key_dict.token
|
||||
|
||||
|
||||
# Status transitions driven by event_type
|
||||
_EVENT_STATUS_MAP: Dict[str, str] = {
|
||||
"step.started": "running",
|
||||
"step.failed": "failed",
|
||||
"hook.waiting": "paused",
|
||||
"hook.received": "running",
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Request / Response models
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class WorkflowRunCreateRequest(BaseModel):
|
||||
workflow_type: str
|
||||
input: Optional[Dict[str, Any]] = None
|
||||
metadata: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
WorkflowRunStatus = Literal["pending", "running", "paused", "completed", "failed"]
|
||||
|
||||
|
||||
class WorkflowRunUpdateRequest(BaseModel):
|
||||
status: Optional[WorkflowRunStatus] = None
|
||||
output: Optional[Dict[str, Any]] = None
|
||||
metadata: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
class WorkflowEventCreateRequest(BaseModel):
|
||||
event_type: str
|
||||
step_name: str
|
||||
data: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
class WorkflowMessageCreateRequest(BaseModel):
|
||||
role: str
|
||||
content: str
|
||||
session_id: Optional[str] = None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def _get_next_sequence_number(prisma_client: Any, run_id: str, table: str) -> int:
|
||||
"""Return MAX(sequence_number) + 1 for the given run, for either events or messages."""
|
||||
if table == "events":
|
||||
rows = await prisma_client.db.litellm_workflowevent.find_many(
|
||||
where={"run_id": run_id},
|
||||
order={"sequence_number": "desc"},
|
||||
take=1,
|
||||
)
|
||||
else:
|
||||
rows = await prisma_client.db.litellm_workflowmessage.find_many(
|
||||
where={"run_id": run_id},
|
||||
order={"sequence_number": "desc"},
|
||||
take=1,
|
||||
)
|
||||
return (rows[0].sequence_number + 1) if rows else 0
|
||||
|
||||
|
||||
async def _require_run(
|
||||
prisma_client: Any,
|
||||
run_id: str,
|
||||
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
|
||||
) -> Any:
|
||||
"""Return the run or raise 404. For non-admin callers, also enforce key ownership."""
|
||||
run = await prisma_client.db.litellm_workflowrun.find_unique(
|
||||
where={"run_id": run_id}
|
||||
)
|
||||
if run is None:
|
||||
raise HTTPException(status_code=404, detail=f"Run '{run_id}' not found")
|
||||
if user_api_key_dict is not None and not _is_admin(user_api_key_dict):
|
||||
caller = _caller_key(user_api_key_dict)
|
||||
if not caller or run.created_by != caller:
|
||||
raise HTTPException(status_code=404, detail=f"Run '{run_id}' not found")
|
||||
return run
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Endpoints
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/workflows/runs",
|
||||
tags=["workflow management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def create_workflow_run(
|
||||
data: WorkflowRunCreateRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""Create a new workflow run. Returns run_id and session_id.
|
||||
|
||||
The caller's API key token is stored as created_by so that non-admin keys
|
||||
can only see and modify their own runs.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
|
||||
try:
|
||||
create_data: Dict[str, Any] = {
|
||||
"workflow_type": data.workflow_type,
|
||||
"created_by": _caller_key(user_api_key_dict),
|
||||
}
|
||||
if data.input is not None:
|
||||
create_data["input"] = _json(data.input)
|
||||
if data.metadata is not None:
|
||||
create_data["metadata"] = _json(data.metadata)
|
||||
run = await prisma_client.db.litellm_workflowrun.create(data=create_data)
|
||||
return run
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error creating workflow run: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.get(
|
||||
"/v1/workflows/runs",
|
||||
tags=["workflow management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def list_workflow_runs(
|
||||
workflow_type: Optional[str] = Query(None),
|
||||
status: Optional[str] = Query(None),
|
||||
limit: int = Query(50, ge=1, le=250),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""List workflow runs. Filter by workflow_type and/or status.
|
||||
|
||||
Non-admin callers only see runs created by their own API key.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
|
||||
where: Dict[str, Any] = {}
|
||||
if workflow_type:
|
||||
where["workflow_type"] = workflow_type
|
||||
if status:
|
||||
statuses = [s.strip() for s in status.split(",")]
|
||||
where["status"] = {"in": statuses} if len(statuses) > 1 else statuses[0]
|
||||
|
||||
# Non-admin callers are scoped to their own key.
|
||||
if not _is_admin(user_api_key_dict):
|
||||
caller = _caller_key(user_api_key_dict)
|
||||
if caller:
|
||||
where["created_by"] = caller
|
||||
|
||||
try:
|
||||
runs = await prisma_client.db.litellm_workflowrun.find_many(
|
||||
where=where,
|
||||
order={"created_at": "desc"},
|
||||
take=limit,
|
||||
)
|
||||
return {"runs": runs, "count": len(runs)}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error listing workflow runs: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.get(
|
||||
"/v1/workflows/runs/{run_id}",
|
||||
tags=["workflow management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def get_workflow_run(
|
||||
run_id: str,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""Get a workflow run with its most recent event."""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
|
||||
try:
|
||||
run = await prisma_client.db.litellm_workflowrun.find_unique(
|
||||
where={"run_id": run_id},
|
||||
include={"events": {"order_by": {"sequence_number": "desc"}, "take": 1}},
|
||||
)
|
||||
if run is None:
|
||||
raise HTTPException(status_code=404, detail=f"Run '{run_id}' not found")
|
||||
if not _is_admin(user_api_key_dict):
|
||||
caller = _caller_key(user_api_key_dict)
|
||||
if not caller or run.created_by != caller:
|
||||
raise HTTPException(status_code=404, detail=f"Run '{run_id}' not found")
|
||||
return run
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error getting workflow run: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.patch(
|
||||
"/v1/workflows/runs/{run_id}",
|
||||
tags=["workflow management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def update_workflow_run(
|
||||
run_id: str,
|
||||
data: WorkflowRunUpdateRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""Update status, metadata, or output on a workflow run."""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
|
||||
update: Dict[str, Any] = {}
|
||||
if data.status is not None:
|
||||
update["status"] = data.status
|
||||
if data.output is not None:
|
||||
update["output"] = _json(data.output)
|
||||
if data.metadata is not None:
|
||||
update["metadata"] = _json(data.metadata)
|
||||
|
||||
if not update:
|
||||
raise HTTPException(status_code=400, detail="No fields to update")
|
||||
|
||||
# Enforce ownership before writing.
|
||||
await _require_run(prisma_client, run_id, user_api_key_dict)
|
||||
|
||||
try:
|
||||
run = await prisma_client.db.litellm_workflowrun.update(
|
||||
where={"run_id": run_id},
|
||||
data=update,
|
||||
)
|
||||
if run is None:
|
||||
raise HTTPException(status_code=404, detail=f"Run '{run_id}' not found")
|
||||
return run
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error updating workflow run: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/workflows/runs/{run_id}/events",
|
||||
tags=["workflow management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def append_workflow_event(
|
||||
run_id: str,
|
||||
data: WorkflowEventCreateRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""Append an event to the run's event log. Also updates run.status if event_type maps to a status.
|
||||
|
||||
Sequence numbers use optimistic concurrency: on a unique-constraint collision
|
||||
(concurrent append), retries up to _MAX_SEQUENCE_RETRIES times with a fresh MAX+1.
|
||||
The event+status update is atomic in a single DB transaction.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
|
||||
await _require_run(prisma_client, run_id, user_api_key_dict)
|
||||
|
||||
new_status = _EVENT_STATUS_MAP.get(data.event_type)
|
||||
|
||||
for attempt in range(_MAX_SEQUENCE_RETRIES):
|
||||
try:
|
||||
seq = await _get_next_sequence_number(prisma_client, run_id, "events")
|
||||
event_data: Dict[str, Any] = {
|
||||
"run_id": run_id,
|
||||
"event_type": data.event_type,
|
||||
"step_name": data.step_name,
|
||||
"sequence_number": seq,
|
||||
}
|
||||
if data.data is not None:
|
||||
event_data["data"] = _json(data.data)
|
||||
|
||||
async with prisma_client.db.tx() as tx:
|
||||
event = await tx.litellm_workflowevent.create(data=event_data)
|
||||
if new_status:
|
||||
await tx.litellm_workflowrun.update(
|
||||
where={"run_id": run_id},
|
||||
data={"status": new_status},
|
||||
)
|
||||
|
||||
return event
|
||||
|
||||
except Exception as e:
|
||||
if UniqueViolationError is not None and isinstance(e, UniqueViolationError):
|
||||
if attempt == _MAX_SEQUENCE_RETRIES - 1:
|
||||
verbose_proxy_logger.exception(
|
||||
"Sequence number collision after %d retries for run %s",
|
||||
_MAX_SEQUENCE_RETRIES,
|
||||
run_id,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail="Concurrent write conflict — please retry",
|
||||
)
|
||||
continue
|
||||
verbose_proxy_logger.exception("Error appending workflow event: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
raise HTTPException(
|
||||
status_code=500, detail="Failed to append event"
|
||||
) # pragma: no cover
|
||||
|
||||
|
||||
@router.get(
|
||||
"/v1/workflows/runs/{run_id}/events",
|
||||
tags=["workflow management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def list_workflow_events(
|
||||
run_id: str,
|
||||
limit: int = Query(100, ge=1, le=500),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""Fetch event log for a run, ordered by sequence_number. Default limit 100, max 500."""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
|
||||
await _require_run(prisma_client, run_id, user_api_key_dict)
|
||||
|
||||
try:
|
||||
events = await prisma_client.db.litellm_workflowevent.find_many(
|
||||
where={"run_id": run_id},
|
||||
order={"sequence_number": "asc"},
|
||||
take=limit,
|
||||
)
|
||||
return {"events": events, "count": len(events)}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error listing workflow events: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/workflows/runs/{run_id}/messages",
|
||||
tags=["workflow management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def append_workflow_message(
|
||||
run_id: str,
|
||||
data: WorkflowMessageCreateRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""Append a conversation message. Stores full content (not truncated).
|
||||
|
||||
Uses optimistic concurrency for sequence numbers.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
|
||||
await _require_run(prisma_client, run_id, user_api_key_dict)
|
||||
|
||||
for attempt in range(_MAX_SEQUENCE_RETRIES):
|
||||
try:
|
||||
seq = await _get_next_sequence_number(prisma_client, run_id, "messages")
|
||||
msg_data: Dict[str, Any] = {
|
||||
"run_id": run_id,
|
||||
"role": data.role,
|
||||
"content": data.content,
|
||||
"sequence_number": seq,
|
||||
}
|
||||
if data.session_id is not None:
|
||||
msg_data["session_id"] = data.session_id
|
||||
msg = await prisma_client.db.litellm_workflowmessage.create(data=msg_data)
|
||||
return msg
|
||||
|
||||
except Exception as e:
|
||||
if UniqueViolationError is not None and isinstance(e, UniqueViolationError):
|
||||
if attempt == _MAX_SEQUENCE_RETRIES - 1:
|
||||
verbose_proxy_logger.exception(
|
||||
"Sequence number collision after %d retries for run %s",
|
||||
_MAX_SEQUENCE_RETRIES,
|
||||
run_id,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
detail="Concurrent write conflict — please retry",
|
||||
)
|
||||
continue
|
||||
verbose_proxy_logger.exception("Error appending workflow message: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
raise HTTPException(
|
||||
status_code=500, detail="Failed to append message"
|
||||
) # pragma: no cover
|
||||
|
||||
|
||||
@router.get(
|
||||
"/v1/workflows/runs/{run_id}/messages",
|
||||
tags=["workflow management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def list_workflow_messages(
|
||||
run_id: str,
|
||||
limit: int = Query(100, ge=1, le=500),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""Fetch conversation history for a run, ordered by sequence_number. Default limit 100, max 500."""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500, detail=CommonProxyErrors.db_not_connected_error.value
|
||||
)
|
||||
|
||||
await _require_run(prisma_client, run_id, user_api_key_dict)
|
||||
|
||||
try:
|
||||
messages = await prisma_client.db.litellm_workflowmessage.find_many(
|
||||
where={"run_id": run_id},
|
||||
order={"sequence_number": "asc"},
|
||||
take=limit,
|
||||
)
|
||||
return {"messages": messages, "count": len(messages)}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error listing workflow messages: %s", e)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
|
@ -426,6 +426,9 @@ from litellm.proxy.management_endpoints.team_endpoints import (
|
|||
from litellm.proxy.management_endpoints.tool_management_endpoints import (
|
||||
router as tool_management_router,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.workflow_management_endpoints import (
|
||||
router as workflow_management_router,
|
||||
)
|
||||
from litellm.proxy.memory.memory_endpoints import router as memory_router
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
get_disabled_non_admin_personal_key_creation,
|
||||
|
|
@ -745,6 +748,10 @@ async def _initialize_shared_aiohttp_session():
|
|||
try:
|
||||
from aiohttp import ClientSession, TCPConnector
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
_build_aiohttp_keepalive_socket_factory,
|
||||
)
|
||||
|
||||
connector_kwargs: Dict[str, Any] = {
|
||||
"keepalive_timeout": AIOHTTP_KEEPALIVE_TIMEOUT,
|
||||
"ttl_dns_cache": AIOHTTP_TTL_DNS_CACHE,
|
||||
|
|
@ -755,6 +762,9 @@ async def _initialize_shared_aiohttp_session():
|
|||
connector_kwargs["limit"] = AIOHTTP_CONNECTOR_LIMIT
|
||||
if AIOHTTP_CONNECTOR_LIMIT_PER_HOST > 0:
|
||||
connector_kwargs["limit_per_host"] = AIOHTTP_CONNECTOR_LIMIT_PER_HOST
|
||||
socket_factory = _build_aiohttp_keepalive_socket_factory()
|
||||
if socket_factory is not None:
|
||||
connector_kwargs["socket_factory"] = socket_factory
|
||||
|
||||
connector = TCPConnector(**connector_kwargs)
|
||||
session = ClientSession(connector=connector)
|
||||
|
|
@ -14279,6 +14289,7 @@ app.include_router(model_management_router)
|
|||
app.include_router(model_access_group_management_router)
|
||||
app.include_router(tag_management_router)
|
||||
app.include_router(tool_management_router)
|
||||
app.include_router(workflow_management_router)
|
||||
app.include_router(memory_router)
|
||||
app.include_router(scheduled_tasks_router)
|
||||
app.include_router(cost_tracking_settings_router)
|
||||
|
|
|
|||
|
|
@ -1291,7 +1291,6 @@ model LiteLLM_AdaptiveRouterSession {
|
|||
@@index([last_activity_at], map: "idx_adaptive_router_session_activity")
|
||||
}
|
||||
|
||||
|
||||
model LiteLLM_ScheduledTaskTable {
|
||||
task_id String @id @default(uuid())
|
||||
// owner_token holds the hashed verification token of the calling key.
|
||||
|
|
@ -1332,3 +1331,80 @@ model LiteLLM_ScheduledTaskTable {
|
|||
@@index([team_id])
|
||||
@@index([agent_id, status])
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Workflow Run Tracking
|
||||
//
|
||||
// Generic durable state tracking for any agent or automated workflow.
|
||||
// Design: three tables — run (header + materialized status), event (append-only
|
||||
// source of truth for state transitions), message (conversation inbox/outbox).
|
||||
//
|
||||
// Usage:
|
||||
// - Set `workflow_type` to identify the owning system (e.g. "shin-builder").
|
||||
// - Store domain-specific fields in `metadata` (worktree_path, pr_url, etc.).
|
||||
// - `session_id` on WorkflowRun matches `x-litellm-session-id` header sent to
|
||||
// the proxy — all spend logs for this run are automatically tagged.
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// One instance of work being done. `status` is a materialized cache of the
|
||||
// latest event; the event log is the authoritative source of truth.
|
||||
model LiteLLM_WorkflowRun {
|
||||
run_id String @id @default(uuid())
|
||||
session_id String @unique @default(uuid())
|
||||
workflow_type String
|
||||
status String @default("pending")
|
||||
created_by String? // user_id of the key that created this run; null = created by master key
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
input Json?
|
||||
output Json?
|
||||
metadata Json?
|
||||
|
||||
events LiteLLM_WorkflowEvent[]
|
||||
messages LiteLLM_WorkflowMessage[]
|
||||
|
||||
@@index([workflow_type, status])
|
||||
@@index([session_id])
|
||||
@@index([created_at])
|
||||
@@index([created_by])
|
||||
}
|
||||
|
||||
// Append-only log of state transitions. Never mutate rows here.
|
||||
// `step_name` and `event_type` are caller-defined strings — no hardcoded enums.
|
||||
// Status auto-update rules (applied by the append endpoint):
|
||||
// step.started → run.status = running
|
||||
// step.failed → run.status = failed
|
||||
// hook.waiting → run.status = paused
|
||||
// hook.received → run.status = running
|
||||
model LiteLLM_WorkflowEvent {
|
||||
event_id String @id @default(uuid())
|
||||
run_id String
|
||||
event_type String
|
||||
step_name String
|
||||
sequence_number Int
|
||||
data Json?
|
||||
created_at DateTime @default(now())
|
||||
|
||||
run LiteLLM_WorkflowRun @relation(fields: [run_id], references: [run_id])
|
||||
|
||||
@@unique([run_id, sequence_number])
|
||||
@@index([run_id])
|
||||
}
|
||||
|
||||
// Conversation inbox/outbox — full message content, separate from the durable
|
||||
// event log. Spend logs truncate messages; this table stores them in full.
|
||||
// `session_id` here is the Claude --resume session ID (or similar).
|
||||
model LiteLLM_WorkflowMessage {
|
||||
message_id String @id @default(uuid())
|
||||
run_id String
|
||||
role String
|
||||
content String
|
||||
sequence_number Int
|
||||
session_id String?
|
||||
created_at DateTime @default(now())
|
||||
|
||||
run LiteLLM_WorkflowRun @relation(fields: [run_id], references: [run_id])
|
||||
|
||||
@@unique([run_id, sequence_number])
|
||||
@@index([run_id])
|
||||
}
|
||||
|
|
|
|||
|
|
@ -106,7 +106,10 @@ from litellm.proxy.db.create_views import (
|
|||
should_create_missing_views,
|
||||
)
|
||||
from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.proxy.db.exception_handler import (
|
||||
PrismaDBExceptionHandler,
|
||||
call_with_db_reconnect_retry,
|
||||
)
|
||||
from litellm.proxy.db.log_db_metrics import log_db_metrics
|
||||
from litellm.proxy.db.prisma_client import PrismaWrapper
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
|
|
@ -2779,30 +2782,42 @@ class PrismaClient:
|
|||
table_name: Literal["users", "keys", "config", "spend"],
|
||||
):
|
||||
"""
|
||||
Generic implementation of get data
|
||||
Generic implementation of get data.
|
||||
|
||||
Self-heals across a single transient transport blip via
|
||||
`call_with_db_reconnect_retry`: on `httpx.ReadError` /
|
||||
`ClientNotConnectedError` / similar, attempt one DB reconnect and
|
||||
retry once before surfacing the failure. Restores the 1.82.6 behavior
|
||||
that was lost in 1.83.x — see issue #25143.
|
||||
"""
|
||||
start_time = time.time()
|
||||
try:
|
||||
|
||||
async def _do_query():
|
||||
if table_name == "users":
|
||||
response = await self.db.litellm_usertable.find_first(
|
||||
return await self.db.litellm_usertable.find_first(
|
||||
where={key: value} # type: ignore
|
||||
)
|
||||
elif table_name == "keys":
|
||||
response = await self.db.litellm_verificationtoken.find_first( # type: ignore
|
||||
return await self.db.litellm_verificationtoken.find_first( # type: ignore
|
||||
where={key: value} # type: ignore
|
||||
)
|
||||
elif table_name == "config":
|
||||
response = await self.db.litellm_config.find_first( # type: ignore
|
||||
return await self.db.litellm_config.find_first( # type: ignore
|
||||
where={key: value} # type: ignore
|
||||
)
|
||||
elif table_name == "spend":
|
||||
response = await self.db.l.find_first( # type: ignore
|
||||
return await self.db.l.find_first( # type: ignore
|
||||
where={key: value} # type: ignore
|
||||
)
|
||||
return response
|
||||
except Exception as e:
|
||||
import traceback
|
||||
return None
|
||||
|
||||
try:
|
||||
return await call_with_db_reconnect_retry(
|
||||
self,
|
||||
_do_query,
|
||||
reason=f"prisma_get_generic_data_{table_name}_lookup_failure",
|
||||
)
|
||||
except Exception as e:
|
||||
error_msg = f"LiteLLM Prisma Client Exception get_generic_data: {str(e)}"
|
||||
verbose_proxy_logger.error(error_msg)
|
||||
error_msg = error_msg + "\nException Type: {}".format(type(e))
|
||||
|
|
@ -4183,8 +4198,11 @@ class PrismaClient:
|
|||
|
||||
Uses the _engine_confirmed_dead flag (set by waitpid thread / pidfd / poll
|
||||
handlers) to choose between heavy reconnect (engine dead -- recreate
|
||||
Prisma client, re-arm watcher) and lightweight reconnect (network
|
||||
blip -- disconnect, connect, SELECT 1).
|
||||
Prisma client, re-arm watcher) and direct reconnect (network blip --
|
||||
recreate Prisma client, re-arm watcher, SELECT 1). Both paths recreate
|
||||
the client via the non-blocking kill-then-construct flow rather than
|
||||
calling disconnect(), which blocks the event loop on the synchronous
|
||||
subprocess.Popen.wait() inside prisma-client-py (see issue #26191).
|
||||
"""
|
||||
effective_timeout = (
|
||||
timeout_seconds
|
||||
|
|
@ -4204,7 +4222,6 @@ class PrismaClient:
|
|||
)
|
||||
self._reap_all_zombies()
|
||||
self._cleanup_engine_watcher()
|
||||
self._engine_confirmed_dead = False
|
||||
|
||||
async def _do_heavy_reconnect() -> None:
|
||||
db_url = os.getenv("DATABASE_URL", "")
|
||||
|
|
@ -4217,23 +4234,32 @@ class PrismaClient:
|
|||
await self._start_engine_watcher()
|
||||
|
||||
await asyncio.wait_for(_do_heavy_reconnect(), timeout=effective_timeout)
|
||||
# Only clear the "dead engine" flag after the heavy reconnect
|
||||
# actually completed. If `_do_heavy_reconnect()` raises (timeout,
|
||||
# missing DATABASE_URL, recreate failure), the flag stays True so
|
||||
# the next attempt re-enters the heavy branch instead of silently
|
||||
# demoting to the lightweight path.
|
||||
self._engine_confirmed_dead = False
|
||||
else:
|
||||
verbose_proxy_logger.debug(
|
||||
"Performing Prisma DB reconnect (engine alive or unknown)."
|
||||
)
|
||||
|
||||
async def _do_direct_reconnect() -> None:
|
||||
old_pid = self._get_engine_pid()
|
||||
try:
|
||||
await self.db.disconnect()
|
||||
except Exception as disconnect_err:
|
||||
verbose_proxy_logger.warning(
|
||||
"Prisma DB disconnect before reconnect failed: %s",
|
||||
disconnect_err,
|
||||
db_url = os.getenv("DATABASE_URL", "")
|
||||
if not db_url:
|
||||
verbose_proxy_logger.error(
|
||||
"DATABASE_URL not set; cannot reconnect Prisma client."
|
||||
)
|
||||
await PrismaWrapper._kill_engine_process(old_pid)
|
||||
|
||||
await self.db.connect()
|
||||
raise RuntimeError("DATABASE_URL not set")
|
||||
# Fresh Prisma client + new engine subprocess. The previous
|
||||
# "lightweight" path called `disconnect()` which blocks the
|
||||
# event loop on `subprocess.Popen.wait()`; since that call
|
||||
# ends up killing the engine anyway, we do it non-blockingly
|
||||
# via `_kill_engine_process` inside `recreate_prisma_client`.
|
||||
self._cleanup_engine_watcher()
|
||||
await self.db.recreate_prisma_client(db_url)
|
||||
await self._start_engine_watcher()
|
||||
await self.db.query_raw("SELECT 1")
|
||||
|
||||
await asyncio.wait_for(_do_direct_reconnect(), timeout=effective_timeout)
|
||||
|
|
|
|||
150
litellm/proxy/workflows/README.md
Normal file
150
litellm/proxy/workflows/README.md
Normal file
|
|
@ -0,0 +1,150 @@
|
|||
# Workflow Run Tracking
|
||||
|
||||
Generic durable state tracking for agents and automated workflows built on the LiteLLM proxy.
|
||||
|
||||
## The Problem
|
||||
|
||||
Agents like [shin-builder](https://github.com/BerriAI/shin-builder) run multi-stage pipelines (triage → plan → implement → PR). Their task state and conversation history lived in memory — a process restart lost everything.
|
||||
|
||||
## Three-Table Design
|
||||
|
||||
```
|
||||
WorkflowRun one instance of work (header + materialized status)
|
||||
WorkflowEvent append-only state transitions (source of truth for replay)
|
||||
WorkflowMessage conversation inbox/outbox (full content, not truncated)
|
||||
```
|
||||
|
||||
**WorkflowEvent is the source of truth.** `WorkflowRun.status` is a materialized cache updated automatically when events are appended. If you need to debug a run, replay its events.
|
||||
|
||||
## API
|
||||
|
||||
All endpoints require a valid LiteLLM API key (`Authorization: Bearer sk-...`).
|
||||
|
||||
### Runs
|
||||
|
||||
```
|
||||
POST /v1/workflows/runs Create a run
|
||||
GET /v1/workflows/runs List runs (?workflow_type=&status=)
|
||||
GET /v1/workflows/runs/{run_id} Get run + latest event
|
||||
PATCH /v1/workflows/runs/{run_id} Update status / metadata / output
|
||||
```
|
||||
|
||||
### Events
|
||||
|
||||
```
|
||||
POST /v1/workflows/runs/{run_id}/events Append event (auto-updates run status)
|
||||
GET /v1/workflows/runs/{run_id}/events Full event log (ordered by sequence)
|
||||
```
|
||||
|
||||
### Messages
|
||||
|
||||
```
|
||||
POST /v1/workflows/runs/{run_id}/messages Append message
|
||||
GET /v1/workflows/runs/{run_id}/messages Conversation history (ordered by sequence)
|
||||
```
|
||||
|
||||
## Quick Start
|
||||
|
||||
```bash
|
||||
# Create a run
|
||||
curl -X POST http://localhost:4000/v1/workflows/runs \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"workflow_type": "shin-builder", "metadata": {"title": "Fix login bug"}}'
|
||||
|
||||
# {"run_id": "abc-123", "session_id": "xyz-456", "status": "pending", ...}
|
||||
|
||||
# Mark step started (sets status → running)
|
||||
curl -X POST http://localhost:4000/v1/workflows/runs/abc-123/events \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"event_type": "step.started", "step_name": "grill", "data": {"claude_session_id": "sess-789"}}'
|
||||
|
||||
# Store a conversation message
|
||||
curl -X POST http://localhost:4000/v1/workflows/runs/abc-123/messages \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"role": "user", "content": "What is the expected behavior?", "session_id": "sess-789"}'
|
||||
|
||||
# Restart recovery: fetch active runs and resume from last event's data.claude_session_id
|
||||
curl "http://localhost:4000/v1/workflows/runs?status=running,paused&workflow_type=shin-builder" \
|
||||
-H "Authorization: Bearer sk-1234"
|
||||
```
|
||||
|
||||
## Status Auto-Update Rules
|
||||
|
||||
When you append an event, the run's status is updated automatically:
|
||||
|
||||
| event_type | run.status |
|
||||
|-----------------|------------|
|
||||
| `step.started` | `running` |
|
||||
| `step.failed` | `failed` |
|
||||
| `hook.waiting` | `paused` |
|
||||
| `hook.received` | `running` |
|
||||
|
||||
Set `status = completed` explicitly via PATCH when the workflow finishes.
|
||||
|
||||
## Linking to Spend Logs
|
||||
|
||||
`WorkflowRun.session_id` is generated automatically (UUID). Pass it as the `x-litellm-session-id` header when making completions through the proxy:
|
||||
|
||||
```python
|
||||
headers = {"x-litellm-session-id": run.session_id}
|
||||
```
|
||||
|
||||
All spend log entries for this run are then tagged automatically. Query cost per run:
|
||||
|
||||
```
|
||||
POST /ui/spend_logs/view_session_spend_logs?session_id={run.session_id}
|
||||
```
|
||||
|
||||
## Sequence Numbers
|
||||
|
||||
Sequence numbers on events and messages are assigned server-side (`MAX + 1` per run). Callers never supply them. This guarantees ordering even under concurrent writes.
|
||||
|
||||
## Using from shin-builder
|
||||
|
||||
Replace the in-memory `tasks.py` dict with calls to these endpoints:
|
||||
|
||||
```python
|
||||
import httpx
|
||||
|
||||
class WorkflowRunClient:
|
||||
def __init__(self, base_url: str, api_key: str):
|
||||
self._client = httpx.AsyncClient(
|
||||
base_url=base_url,
|
||||
headers={"Authorization": f"Bearer {api_key}"},
|
||||
)
|
||||
|
||||
async def create_task(self, title: str, **metadata) -> dict:
|
||||
r = await self._client.post("/v1/workflows/runs", json={
|
||||
"workflow_type": "shin-builder",
|
||||
"metadata": {"title": title, **metadata},
|
||||
})
|
||||
r.raise_for_status()
|
||||
return r.json()
|
||||
|
||||
async def list_active_tasks(self) -> list:
|
||||
r = await self._client.get(
|
||||
"/v1/workflows/runs",
|
||||
params={"workflow_type": "shin-builder", "status": "running,paused"},
|
||||
)
|
||||
r.raise_for_status()
|
||||
return r.json()["runs"]
|
||||
|
||||
async def transition(self, run_id: str, step_name: str, event_type: str, data: dict = None):
|
||||
r = await self._client.post(f"/v1/workflows/runs/{run_id}/events", json={
|
||||
"event_type": event_type,
|
||||
"step_name": step_name,
|
||||
"data": data or {},
|
||||
})
|
||||
r.raise_for_status()
|
||||
|
||||
async def append_message(self, run_id: str, role: str, content: str, session_id: str = None):
|
||||
r = await self._client.post(f"/v1/workflows/runs/{run_id}/messages", json={
|
||||
"role": role, "content": content, "session_id": session_id,
|
||||
})
|
||||
r.raise_for_status()
|
||||
```
|
||||
|
||||
On startup, call `list_active_tasks()` to restore in-flight runs. The last `step.started` event's `data.claude_session_id` gives you the `--resume` ID.
|
||||
|
|
@ -81,6 +81,12 @@ class MCPServer(BaseModel):
|
|||
# Defaults to the token's expires_in minus the expiry buffer, or
|
||||
# MCP_PER_USER_TOKEN_DEFAULT_TTL when expires_in is absent.
|
||||
token_storage_ttl_seconds: Optional[int] = None
|
||||
# Resolved short-ID tool prefix when LITELLM_USE_SHORT_MCP_TOOL_PREFIX is
|
||||
# enabled. Set by ``MCPServerManager._assign_unique_short_prefix`` at
|
||||
# registration time so that natural-hash collisions between two
|
||||
# different ``server_id`` values are bumped deterministically. Left
|
||||
# ``None`` in default-prefix mode.
|
||||
short_prefix: Optional[str] = None
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
@property
|
||||
|
|
|
|||
|
|
@ -6526,6 +6526,7 @@ def validate_environment( # noqa: PLR0915
|
|||
or model in litellm.open_ai_text_completion_models
|
||||
or model in litellm.open_ai_embedding_models
|
||||
or model in litellm.openai_image_generation_models
|
||||
or model.startswith("gpt-image")
|
||||
):
|
||||
if "OPENAI_API_KEY" in os.environ:
|
||||
keys_in_environment = True
|
||||
|
|
|
|||
|
|
@ -5117,6 +5117,38 @@
|
|||
"/v1/images/edits"
|
||||
]
|
||||
},
|
||||
"azure/gpt-image-2": {
|
||||
"cache_read_input_image_token_cost": 2e-06,
|
||||
"cache_read_input_token_cost": 1.25e-06,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"input_cost_per_image_token": 8e-06,
|
||||
"litellm_provider": "azure",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_token": 1e-05,
|
||||
"output_cost_per_image_token": 3e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supports_vision": true,
|
||||
"supports_pdf_input": true
|
||||
},
|
||||
"azure/gpt-image-2-2026-04-21": {
|
||||
"cache_read_input_image_token_cost": 2e-06,
|
||||
"cache_read_input_token_cost": 1.25e-06,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"input_cost_per_image_token": 8e-06,
|
||||
"litellm_provider": "azure",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_token": 1e-05,
|
||||
"output_cost_per_image_token": 3e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supports_vision": true,
|
||||
"supports_pdf_input": true
|
||||
},
|
||||
"azure/low/1024-x-1024/gpt-image-1-mini": {
|
||||
"input_cost_per_pixel": 2.0751953125e-09,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -19097,6 +19129,38 @@
|
|||
"supports_vision": true,
|
||||
"supports_pdf_input": true
|
||||
},
|
||||
"gpt-image-2": {
|
||||
"cache_read_input_image_token_cost": 2e-06,
|
||||
"cache_read_input_token_cost": 1.25e-06,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_token": 1e-05,
|
||||
"input_cost_per_image_token": 8e-06,
|
||||
"output_cost_per_image_token": 3e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supports_vision": true,
|
||||
"supports_pdf_input": true
|
||||
},
|
||||
"gpt-image-2-2026-04-21": {
|
||||
"cache_read_input_image_token_cost": 2e-06,
|
||||
"cache_read_input_token_cost": 1.25e-06,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"litellm_provider": "openai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_token": 1e-05,
|
||||
"input_cost_per_image_token": 8e-06,
|
||||
"output_cost_per_image_token": 3e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits"
|
||||
],
|
||||
"supports_vision": true,
|
||||
"supports_pdf_input": true
|
||||
},
|
||||
"low/1024-x-1024/gpt-image-1.5": {
|
||||
"input_cost_per_image": 0.009,
|
||||
"litellm_provider": "openai",
|
||||
|
|
|
|||
|
|
@ -193,6 +193,23 @@
|
|||
"a2a": false
|
||||
}
|
||||
},
|
||||
"aihubmix": {
|
||||
"display_name": "AIHubMix (`aihubmix`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/aihubmix",
|
||||
"endpoints": {
|
||||
"chat_completions": true,
|
||||
"messages": true,
|
||||
"responses": true,
|
||||
"embeddings": true,
|
||||
"image_generations": true,
|
||||
"audio_transcriptions": true,
|
||||
"audio_speech": true,
|
||||
"moderations": true,
|
||||
"batches": false,
|
||||
"rerank": true,
|
||||
"a2a": false
|
||||
}
|
||||
},
|
||||
"assemblyai": {
|
||||
"display_name": "AssemblyAI (`assemblyai`)",
|
||||
"url": "https://docs.litellm.ai/docs/pass_through/assembly_ai",
|
||||
|
|
|
|||
|
|
@ -1291,7 +1291,6 @@ model LiteLLM_AdaptiveRouterSession {
|
|||
@@index([last_activity_at], map: "idx_adaptive_router_session_activity")
|
||||
}
|
||||
|
||||
|
||||
model LiteLLM_ScheduledTaskTable {
|
||||
task_id String @id @default(uuid())
|
||||
// owner_token holds the hashed verification token of the calling key.
|
||||
|
|
@ -1332,3 +1331,80 @@ model LiteLLM_ScheduledTaskTable {
|
|||
@@index([team_id])
|
||||
@@index([agent_id, status])
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Workflow Run Tracking
|
||||
//
|
||||
// Generic durable state tracking for any agent or automated workflow.
|
||||
// Design: three tables — run (header + materialized status), event (append-only
|
||||
// source of truth for state transitions), message (conversation inbox/outbox).
|
||||
//
|
||||
// Usage:
|
||||
// - Set `workflow_type` to identify the owning system (e.g. "shin-builder").
|
||||
// - Store domain-specific fields in `metadata` (worktree_path, pr_url, etc.).
|
||||
// - `session_id` on WorkflowRun matches `x-litellm-session-id` header sent to
|
||||
// the proxy — all spend logs for this run are automatically tagged.
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// One instance of work being done. `status` is a materialized cache of the
|
||||
// latest event; the event log is the authoritative source of truth.
|
||||
model LiteLLM_WorkflowRun {
|
||||
run_id String @id @default(uuid())
|
||||
session_id String @unique @default(uuid())
|
||||
workflow_type String
|
||||
status String @default("pending")
|
||||
created_by String? // user_id of the key that created this run; null = created by master key
|
||||
created_at DateTime @default(now())
|
||||
updated_at DateTime @updatedAt
|
||||
input Json?
|
||||
output Json?
|
||||
metadata Json?
|
||||
|
||||
events LiteLLM_WorkflowEvent[]
|
||||
messages LiteLLM_WorkflowMessage[]
|
||||
|
||||
@@index([workflow_type, status])
|
||||
@@index([session_id])
|
||||
@@index([created_at])
|
||||
@@index([created_by])
|
||||
}
|
||||
|
||||
// Append-only log of state transitions. Never mutate rows here.
|
||||
// `step_name` and `event_type` are caller-defined strings — no hardcoded enums.
|
||||
// Status auto-update rules (applied by the append endpoint):
|
||||
// step.started → run.status = running
|
||||
// step.failed → run.status = failed
|
||||
// hook.waiting → run.status = paused
|
||||
// hook.received → run.status = running
|
||||
model LiteLLM_WorkflowEvent {
|
||||
event_id String @id @default(uuid())
|
||||
run_id String
|
||||
event_type String
|
||||
step_name String
|
||||
sequence_number Int
|
||||
data Json?
|
||||
created_at DateTime @default(now())
|
||||
|
||||
run LiteLLM_WorkflowRun @relation(fields: [run_id], references: [run_id])
|
||||
|
||||
@@unique([run_id, sequence_number])
|
||||
@@index([run_id])
|
||||
}
|
||||
|
||||
// Conversation inbox/outbox — full message content, separate from the durable
|
||||
// event log. Spend logs truncate messages; this table stores them in full.
|
||||
// `session_id` here is the Claude --resume session ID (or similar).
|
||||
model LiteLLM_WorkflowMessage {
|
||||
message_id String @id @default(uuid())
|
||||
run_id String
|
||||
role String
|
||||
content String
|
||||
sequence_number Int
|
||||
session_id String?
|
||||
created_at DateTime @default(now())
|
||||
|
||||
run LiteLLM_WorkflowRun @relation(fields: [run_id], references: [run_id])
|
||||
|
||||
@@unique([run_id, sequence_number])
|
||||
@@index([run_id])
|
||||
}
|
||||
|
|
|
|||
|
|
@ -254,32 +254,48 @@ async def test_run_reconnect_cycle_uses_heavy_path_when_confirmed_dead(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_reconnect_cycle_uses_lightweight_path_when_engine_alive(
|
||||
async def test_run_reconnect_cycle_uses_direct_path_when_engine_alive(
|
||||
engine_client,
|
||||
) -> None:
|
||||
"""_run_reconnect_cycle uses disconnect/connect when engine is alive."""
|
||||
engine_client._engine_pid = 1234
|
||||
"""Direct reconnect (engine alive) calls recreate_prisma_client + SELECT 1.
|
||||
|
||||
with patch.object(engine_client, "_is_engine_alive", return_value=True):
|
||||
The old "lightweight" path called `disconnect()` + `connect()`, which
|
||||
blocks the event loop on the sync `process.wait()` inside aclose().
|
||||
The fix routes both engine-alive and engine-dead paths through
|
||||
`recreate_prisma_client`, which non-blockingly kills the old engine.
|
||||
"""
|
||||
engine_client._engine_pid = 1234
|
||||
engine_client._start_engine_watcher = AsyncMock()
|
||||
|
||||
with (
|
||||
patch.object(engine_client, "_is_engine_alive", return_value=True),
|
||||
patch.dict(os.environ, {"DATABASE_URL": "postgresql://test"}),
|
||||
):
|
||||
await engine_client._run_reconnect_cycle(timeout_seconds=5.0)
|
||||
|
||||
engine_client.db.connect.assert_awaited_once()
|
||||
engine_client.db.recreate_prisma_client.assert_awaited_once_with(
|
||||
"postgresql://test"
|
||||
)
|
||||
engine_client.db.query_raw.assert_awaited_once_with("SELECT 1")
|
||||
engine_client.db.recreate_prisma_client.assert_not_awaited()
|
||||
engine_client.db.disconnect.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_reconnect_cycle_uses_lightweight_path_when_pid_unknown(
|
||||
async def test_run_reconnect_cycle_uses_direct_path_when_pid_unknown(
|
||||
engine_client,
|
||||
) -> None:
|
||||
"""_run_reconnect_cycle uses lightweight path when engine PID is not tracked."""
|
||||
"""When the engine PID is not tracked, direct reconnect still runs."""
|
||||
engine_client._engine_pid = 0
|
||||
engine_client._start_engine_watcher = AsyncMock()
|
||||
|
||||
await engine_client._run_reconnect_cycle(timeout_seconds=5.0)
|
||||
with patch.dict(os.environ, {"DATABASE_URL": "postgresql://test"}):
|
||||
await engine_client._run_reconnect_cycle(timeout_seconds=5.0)
|
||||
|
||||
engine_client.db.connect.assert_awaited_once()
|
||||
engine_client.db.recreate_prisma_client.assert_awaited_once_with(
|
||||
"postgresql://test"
|
||||
)
|
||||
engine_client.db.query_raw.assert_awaited_once_with("SELECT 1")
|
||||
engine_client.db.recreate_prisma_client.assert_not_awaited()
|
||||
engine_client.db.disconnect.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -473,36 +489,38 @@ def test_on_engine_death_from_thread_ignores_stale_pid(engine_client):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_escalation_after_consecutive_lightweight_failures(engine_client):
|
||||
"""After N consecutive lightweight reconnect failures, _engine_confirmed_dead
|
||||
async def test_escalation_after_consecutive_direct_reconnect_failures(engine_client):
|
||||
"""After N consecutive direct reconnect failures, _engine_confirmed_dead
|
||||
is set to True so _run_reconnect_cycle takes the heavy reconnect path."""
|
||||
engine_client._reconnect_escalation_threshold = 3
|
||||
engine_client._consecutive_reconnect_failures = 0
|
||||
engine_client._db_reconnect_cooldown_seconds = 0 # disable cooldown for test
|
||||
engine_client._start_engine_watcher = AsyncMock(return_value=None)
|
||||
|
||||
# Make lightweight reconnect fail every time
|
||||
engine_client.db.disconnect = AsyncMock(return_value=None)
|
||||
engine_client.db.connect = AsyncMock(side_effect=Exception("connect failed"))
|
||||
# Make direct reconnect fail every time
|
||||
engine_client.db.recreate_prisma_client = AsyncMock(
|
||||
side_effect=Exception("recreate failed")
|
||||
)
|
||||
|
||||
# Run 3 failed reconnect attempts
|
||||
for i in range(3):
|
||||
result = await engine_client._attempt_reconnect_inside_lock(
|
||||
force=True, reason="test", timeout_seconds=5.0
|
||||
)
|
||||
assert result is False
|
||||
with patch.dict(os.environ, {"DATABASE_URL": "postgresql://test"}):
|
||||
for _ in range(3):
|
||||
result = await engine_client._attempt_reconnect_inside_lock(
|
||||
force=True, reason="test", timeout_seconds=5.0
|
||||
)
|
||||
assert result is False
|
||||
|
||||
assert engine_client._consecutive_reconnect_failures == 3
|
||||
|
||||
# Next attempt should escalate: _engine_confirmed_dead set to True before _run_reconnect_cycle
|
||||
# Next attempt should escalate to the heavy path (recreate_prisma_client still
|
||||
# the call, but via the _engine_confirmed_dead branch that also re-arms the watcher).
|
||||
engine_client.db.recreate_prisma_client = AsyncMock(return_value=None)
|
||||
engine_client._start_engine_watcher = AsyncMock(return_value=None)
|
||||
|
||||
with patch.dict(os.environ, {"DATABASE_URL": "postgresql://test"}):
|
||||
result = await engine_client._attempt_reconnect_inside_lock(
|
||||
force=True, reason="test_escalation", timeout_seconds=5.0
|
||||
)
|
||||
|
||||
# Heavy reconnect should have been attempted (recreate_prisma_client called)
|
||||
engine_client.db.recreate_prisma_client.assert_awaited_once()
|
||||
|
||||
|
||||
|
|
@ -511,15 +529,16 @@ async def test_successful_reconnect_resets_failure_counter(engine_client):
|
|||
"""A successful reconnect resets _consecutive_reconnect_failures to 0."""
|
||||
engine_client._consecutive_reconnect_failures = 2
|
||||
engine_client._db_reconnect_cooldown_seconds = 0
|
||||
engine_client._start_engine_watcher = AsyncMock()
|
||||
|
||||
# Make reconnect succeed
|
||||
engine_client.db.disconnect = AsyncMock(return_value=None)
|
||||
engine_client.db.connect = AsyncMock(return_value=None)
|
||||
engine_client.db.recreate_prisma_client = AsyncMock(return_value=None)
|
||||
engine_client.db.query_raw = AsyncMock(return_value=[{"result": 1}])
|
||||
|
||||
result = await engine_client._attempt_reconnect_inside_lock(
|
||||
force=True, reason="test", timeout_seconds=5.0
|
||||
)
|
||||
with patch.dict(os.environ, {"DATABASE_URL": "postgresql://test"}):
|
||||
result = await engine_client._attempt_reconnect_inside_lock(
|
||||
force=True, reason="test", timeout_seconds=5.0
|
||||
)
|
||||
|
||||
assert result is True
|
||||
assert engine_client._consecutive_reconnect_failures == 0
|
||||
|
|
|
|||
|
|
@ -30,10 +30,9 @@ def test_model_alias_map(caplog):
|
|||
)
|
||||
print(response.model)
|
||||
|
||||
captured_logs = [rec.levelname for rec in caplog.records]
|
||||
|
||||
for log in captured_logs:
|
||||
assert "ERROR" not in log
|
||||
for rec in caplog.records:
|
||||
if rec.levelname == "ERROR" and rec.name.startswith("LiteLLM"):
|
||||
pytest.fail(f"Unexpected litellm ERROR log: {rec.getMessage()}")
|
||||
|
||||
assert "llama-3.1-8b-instant" in response.model
|
||||
except litellm.ServiceUnavailableError:
|
||||
|
|
|
|||
|
|
@ -354,6 +354,7 @@ async def test_bedrock_kb_request_body_has_transformed_filters(
|
|||
api_base=api_base,
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params=litellm_params_dict,
|
||||
extra_body=None,
|
||||
)
|
||||
)
|
||||
captured_request_body["url"] = url
|
||||
|
|
|
|||
|
|
@ -4146,3 +4146,39 @@ def test_transform_response_finish_reason_stop_when_json_mode_filters_all_tools(
|
|||
|
||||
# finish_reason must be "stop", not "tool_calls"
|
||||
assert result.choices[0].finish_reason == "stop"
|
||||
|
||||
|
||||
def test_transform_response_does_not_leak_body_on_parse_failure():
|
||||
from litellm.llms.bedrock.common_utils import BedrockError
|
||||
|
||||
leaky_body = {"output": {"message": {"content": [{"text": "secret content"}]}}}
|
||||
|
||||
class MockResponse:
|
||||
def json(self):
|
||||
return leaky_body
|
||||
|
||||
@property
|
||||
def text(self):
|
||||
return json.dumps(leaky_body)
|
||||
|
||||
with patch(
|
||||
"litellm.llms.bedrock.chat.converse_transformation.ConverseResponseBlock",
|
||||
side_effect=KeyError("missing required field"),
|
||||
):
|
||||
with pytest.raises(BedrockError) as exc_info:
|
||||
AmazonConverseConfig()._transform_response(
|
||||
model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
|
||||
response=MockResponse(),
|
||||
model_response=ModelResponse(),
|
||||
stream=False,
|
||||
logging_obj=None,
|
||||
optional_params={},
|
||||
api_key=None,
|
||||
data=None,
|
||||
messages=[],
|
||||
encoding=None,
|
||||
)
|
||||
|
||||
msg = str(exc_info.value)
|
||||
assert "secret content" not in msg
|
||||
assert "Error converting to valid response block" in msg
|
||||
|
|
|
|||
|
|
@ -21,7 +21,112 @@ def test_transform_search_request():
|
|||
api_base="https://bedrock-agent-runtime.us-west-2.amazonaws.com/knowledgebases",
|
||||
litellm_logging_obj=mock_log,
|
||||
litellm_params={},
|
||||
extra_body=None,
|
||||
)
|
||||
|
||||
assert url.endswith("/kb123/retrieve")
|
||||
assert body["retrievalQuery"].get("text") == "hello"
|
||||
|
||||
|
||||
def test_transform_search_request_uses_only_retrieval_config_from_extra_body():
|
||||
config = BedrockVectorStoreConfig()
|
||||
mock_log = MagicMock()
|
||||
mock_log.model_call_details = {}
|
||||
|
||||
url, body = config.transform_search_vector_store_request(
|
||||
vector_store_id="kb123",
|
||||
query="hello",
|
||||
vector_store_search_optional_params={},
|
||||
api_base="https://bedrock-agent-runtime.us-west-2.amazonaws.com/knowledgebases",
|
||||
litellm_logging_obj=mock_log,
|
||||
litellm_params={},
|
||||
extra_body={
|
||||
"retrievalConfiguration": {
|
||||
"vectorSearchConfiguration": {
|
||||
"overrideSearchType": "HYBRID",
|
||||
"numberOfResults": 8,
|
||||
}
|
||||
},
|
||||
"unrelatedField": {"should_not": "be_forwarded"},
|
||||
},
|
||||
)
|
||||
|
||||
assert url.endswith("/kb123/retrieve")
|
||||
assert body["retrievalQuery"].get("text") == "hello"
|
||||
assert (
|
||||
body["retrievalConfiguration"]["vectorSearchConfiguration"][
|
||||
"overrideSearchType"
|
||||
]
|
||||
== "HYBRID"
|
||||
)
|
||||
assert "unrelatedField" not in body
|
||||
|
||||
|
||||
def test_transform_search_request_does_not_mutate_extra_body_and_overrides_number_of_results():
|
||||
config = BedrockVectorStoreConfig()
|
||||
mock_log = MagicMock()
|
||||
mock_log.model_call_details = {}
|
||||
extra_body = {
|
||||
"retrievalConfiguration": {
|
||||
"vectorSearchConfiguration": {
|
||||
"overrideSearchType": "HYBRID",
|
||||
"numberOfResults": 8,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
_, body = config.transform_search_vector_store_request(
|
||||
vector_store_id="kb123",
|
||||
query="hello",
|
||||
vector_store_search_optional_params={"max_num_results": 10},
|
||||
api_base="https://bedrock-agent-runtime.us-west-2.amazonaws.com/knowledgebases",
|
||||
litellm_logging_obj=mock_log,
|
||||
litellm_params={},
|
||||
extra_body=extra_body,
|
||||
)
|
||||
|
||||
assert (
|
||||
body["retrievalConfiguration"]["vectorSearchConfiguration"]["numberOfResults"]
|
||||
== 10
|
||||
)
|
||||
assert (
|
||||
extra_body["retrievalConfiguration"]["vectorSearchConfiguration"][
|
||||
"numberOfResults"
|
||||
]
|
||||
== 8
|
||||
)
|
||||
|
||||
|
||||
def test_transform_search_request_overrides_filter_without_mutating_extra_body():
|
||||
config = BedrockVectorStoreConfig()
|
||||
mock_log = MagicMock()
|
||||
mock_log.model_call_details = {}
|
||||
extra_body = {
|
||||
"retrievalConfiguration": {
|
||||
"vectorSearchConfiguration": {
|
||||
"filter": {"equals": {"key": "tenant", "value": "a"}}
|
||||
}
|
||||
}
|
||||
}
|
||||
new_filter = {"equals": {"key": "tenant", "value": "b"}}
|
||||
|
||||
_, body = config.transform_search_vector_store_request(
|
||||
vector_store_id="kb123",
|
||||
query="hello",
|
||||
vector_store_search_optional_params={"filters": new_filter},
|
||||
api_base="https://bedrock-agent-runtime.us-west-2.amazonaws.com/knowledgebases",
|
||||
litellm_logging_obj=mock_log,
|
||||
litellm_params={},
|
||||
extra_body=extra_body,
|
||||
)
|
||||
|
||||
assert (
|
||||
body["retrievalConfiguration"]["vectorSearchConfiguration"]["filter"]
|
||||
== new_filter
|
||||
)
|
||||
assert (
|
||||
extra_body["retrievalConfiguration"]["vectorSearchConfiguration"]["filter"][
|
||||
"equals"
|
||||
]["value"]
|
||||
== "a"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,161 @@
|
|||
import socket
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
|
||||
def _invoke_connector_factory(http_handler_module):
|
||||
"""
|
||||
Drive the lambda factory installed on the transport so TCPConnector is
|
||||
actually constructed. _create_aiohttp_transport returns a transport whose
|
||||
_client_factory is the lambda that builds (TCPConnector → ClientSession);
|
||||
invoking it directly avoids relying on _get_valid_client_session's internal
|
||||
branching to trigger connector construction.
|
||||
"""
|
||||
transport = http_handler_module.AsyncHTTPHandler._create_aiohttp_transport(
|
||||
shared_session=None
|
||||
)
|
||||
transport._client_factory()
|
||||
return transport
|
||||
|
||||
|
||||
def test_socket_factory_omitted_when_disabled(monkeypatch):
|
||||
from litellm.llms.custom_httpx import http_handler as http_handler_module
|
||||
|
||||
monkeypatch.setattr(http_handler_module, "AIOHTTP_SO_KEEPALIVE", False)
|
||||
monkeypatch.setattr(http_handler_module, "_AIOHTTP_SUPPORTS_SOCKET_FACTORY", True)
|
||||
|
||||
connector_mock = MagicMock(name="connector")
|
||||
session_mock = MagicMock(name="session")
|
||||
|
||||
with patch.object(
|
||||
http_handler_module, "TCPConnector", return_value=connector_mock
|
||||
) as mock_tcp_connector:
|
||||
with patch.object(
|
||||
http_handler_module, "ClientSession", return_value=session_mock
|
||||
):
|
||||
_invoke_connector_factory(http_handler_module)
|
||||
|
||||
assert mock_tcp_connector.call_count >= 1
|
||||
assert "socket_factory" not in mock_tcp_connector.call_args.kwargs
|
||||
|
||||
|
||||
def test_socket_factory_attached_when_enabled(monkeypatch):
|
||||
from litellm.llms.custom_httpx import http_handler as http_handler_module
|
||||
|
||||
monkeypatch.setattr(http_handler_module, "AIOHTTP_SO_KEEPALIVE", True)
|
||||
monkeypatch.setattr(http_handler_module, "_AIOHTTP_SUPPORTS_SOCKET_FACTORY", True)
|
||||
|
||||
connector_mock = MagicMock(name="connector")
|
||||
session_mock = MagicMock(name="session")
|
||||
|
||||
with patch.object(
|
||||
http_handler_module, "TCPConnector", return_value=connector_mock
|
||||
) as mock_tcp_connector:
|
||||
with patch.object(
|
||||
http_handler_module, "ClientSession", return_value=session_mock
|
||||
):
|
||||
_invoke_connector_factory(http_handler_module)
|
||||
|
||||
assert mock_tcp_connector.call_count >= 1
|
||||
factory = mock_tcp_connector.call_args.kwargs.get("socket_factory")
|
||||
assert callable(factory)
|
||||
|
||||
|
||||
def test_socket_factory_skipped_on_old_aiohttp(monkeypatch):
|
||||
from litellm.llms.custom_httpx import http_handler as http_handler_module
|
||||
|
||||
monkeypatch.setattr(http_handler_module, "AIOHTTP_SO_KEEPALIVE", True)
|
||||
monkeypatch.setattr(http_handler_module, "_AIOHTTP_SUPPORTS_SOCKET_FACTORY", False)
|
||||
|
||||
connector_mock = MagicMock(name="connector")
|
||||
session_mock = MagicMock(name="session")
|
||||
|
||||
with patch.object(
|
||||
http_handler_module, "TCPConnector", return_value=connector_mock
|
||||
) as mock_tcp_connector:
|
||||
with patch.object(
|
||||
http_handler_module, "ClientSession", return_value=session_mock
|
||||
):
|
||||
_invoke_connector_factory(http_handler_module)
|
||||
|
||||
assert mock_tcp_connector.call_count >= 1
|
||||
assert "socket_factory" not in mock_tcp_connector.call_args.kwargs
|
||||
|
||||
|
||||
def test_socket_factory_sets_keepalive_options(monkeypatch):
|
||||
from litellm.llms.custom_httpx import http_handler as http_handler_module
|
||||
|
||||
monkeypatch.setattr(http_handler_module, "AIOHTTP_SO_KEEPALIVE", True)
|
||||
monkeypatch.setattr(http_handler_module, "_AIOHTTP_SUPPORTS_SOCKET_FACTORY", True)
|
||||
monkeypatch.setattr(http_handler_module, "AIOHTTP_TCP_KEEPIDLE", 45)
|
||||
monkeypatch.setattr(http_handler_module, "AIOHTTP_TCP_KEEPINTVL", 15)
|
||||
monkeypatch.setattr(http_handler_module, "AIOHTTP_TCP_KEEPCNT", 4)
|
||||
|
||||
factory = http_handler_module._build_aiohttp_keepalive_socket_factory()
|
||||
assert factory is not None
|
||||
|
||||
addr_info = (socket.AF_INET, socket.SOCK_STREAM, socket.IPPROTO_TCP, "", ("", 0))
|
||||
|
||||
fake_sock = MagicMock(spec=socket.socket)
|
||||
with patch("socket.socket", return_value=fake_sock) as sock_ctor:
|
||||
returned = factory(addr_info)
|
||||
|
||||
sock_ctor.assert_called_once_with(
|
||||
family=socket.AF_INET, type=socket.SOCK_STREAM, proto=socket.IPPROTO_TCP
|
||||
)
|
||||
assert returned is fake_sock
|
||||
fake_sock.setblocking.assert_called_once_with(False)
|
||||
|
||||
setsockopt_calls = {
|
||||
(call.args[0], call.args[1]): call.args[2]
|
||||
for call in fake_sock.setsockopt.call_args_list
|
||||
}
|
||||
assert setsockopt_calls[(socket.SOL_SOCKET, socket.SO_KEEPALIVE)] == 1
|
||||
|
||||
if hasattr(socket, "TCP_KEEPIDLE"):
|
||||
assert setsockopt_calls[(socket.IPPROTO_TCP, socket.TCP_KEEPIDLE)] == 45
|
||||
elif hasattr(socket, "TCP_KEEPALIVE"):
|
||||
assert setsockopt_calls[(socket.IPPROTO_TCP, socket.TCP_KEEPALIVE)] == 45
|
||||
if hasattr(socket, "TCP_KEEPINTVL"):
|
||||
assert setsockopt_calls[(socket.IPPROTO_TCP, socket.TCP_KEEPINTVL)] == 15
|
||||
if hasattr(socket, "TCP_KEEPCNT"):
|
||||
assert setsockopt_calls[(socket.IPPROTO_TCP, socket.TCP_KEEPCNT)] == 4
|
||||
|
||||
|
||||
def test_socket_factory_uses_tcp_keepalive_when_keepidle_unavailable(monkeypatch):
|
||||
"""
|
||||
Cover the macOS/Darwin branch: when TCP_KEEPIDLE is missing but TCP_KEEPALIVE
|
||||
is present, the factory should fall back to TCP_KEEPALIVE for the idle timer.
|
||||
Linux CI runners always have TCP_KEEPIDLE, so we patch socket itself to
|
||||
simulate the BSD-derived environment.
|
||||
"""
|
||||
from litellm.llms.custom_httpx import http_handler as http_handler_module
|
||||
|
||||
monkeypatch.setattr(http_handler_module, "AIOHTTP_SO_KEEPALIVE", True)
|
||||
monkeypatch.setattr(http_handler_module, "_AIOHTTP_SUPPORTS_SOCKET_FACTORY", True)
|
||||
monkeypatch.setattr(http_handler_module, "AIOHTTP_TCP_KEEPIDLE", 60)
|
||||
|
||||
factory = http_handler_module._build_aiohttp_keepalive_socket_factory()
|
||||
assert factory is not None
|
||||
|
||||
fake_socket_module = MagicMock(spec=[])
|
||||
fake_socket_module.SOL_SOCKET = socket.SOL_SOCKET
|
||||
fake_socket_module.SO_KEEPALIVE = socket.SO_KEEPALIVE
|
||||
fake_socket_module.IPPROTO_TCP = socket.IPPROTO_TCP
|
||||
fake_socket_module.TCP_KEEPALIVE = getattr(socket, "TCP_KEEPALIVE", 0x10)
|
||||
fake_sock = MagicMock(spec=socket.socket)
|
||||
fake_socket_module.socket = MagicMock(return_value=fake_sock)
|
||||
|
||||
addr_info = (socket.AF_INET, socket.SOCK_STREAM, socket.IPPROTO_TCP, "", ("", 0))
|
||||
|
||||
with patch.object(http_handler_module, "socket", fake_socket_module):
|
||||
factory(addr_info)
|
||||
|
||||
setsockopt_calls = {
|
||||
(call.args[0], call.args[1]): call.args[2]
|
||||
for call in fake_sock.setsockopt.call_args_list
|
||||
}
|
||||
assert setsockopt_calls[(socket.SOL_SOCKET, socket.SO_KEEPALIVE)] == 1
|
||||
assert (
|
||||
setsockopt_calls[(socket.IPPROTO_TCP, fake_socket_module.TCP_KEEPALIVE)] == 60
|
||||
)
|
||||
assert (socket.IPPROTO_TCP, getattr(socket, "TCP_KEEPIDLE", -1)) not in setsockopt_calls
|
||||
|
|
@ -2,12 +2,16 @@ import os
|
|||
import sys
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import (
|
||||
BaseLLMHTTPHandler,
|
||||
_google_genai_streaming_hidden_params,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
|
||||
|
|
@ -320,3 +324,29 @@ async def test_async_anthropic_messages_handler_header_priority():
|
|||
assert captured_headers["X-Forwarded-Only"] == "keep"
|
||||
assert captured_headers["X-Extra-Only"] == "also-keep"
|
||||
assert captured_headers["X-Provider-Only"] == "keep-this-too"
|
||||
|
||||
|
||||
def test_google_genai_streaming_hidden_params_model_info_and_router_fallback():
|
||||
logging_obj = Mock()
|
||||
logging_obj.get_router_model_id = Mock(return_value="router-model-id")
|
||||
|
||||
from_model_info = _google_genai_streaming_hidden_params(
|
||||
api_base="https://generativelanguage.googleapis.com/v1beta",
|
||||
litellm_params=GenericLiteLLMParams(model_info={"id": "info-id"}),
|
||||
logging_obj=logging_obj,
|
||||
response_headers=httpx.Headers({"x-ratelimit-remaining": "10"}),
|
||||
)
|
||||
assert from_model_info["model_id"] == "info-id"
|
||||
assert (
|
||||
from_model_info["api_base"]
|
||||
== "https://generativelanguage.googleapis.com/v1beta"
|
||||
)
|
||||
assert isinstance(from_model_info["additional_headers"], dict)
|
||||
|
||||
from_router = _google_genai_streaming_hidden_params(
|
||||
api_base="https://x",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
logging_obj=logging_obj,
|
||||
response_headers=httpx.Headers({}),
|
||||
)
|
||||
assert from_router["model_id"] == "router-model-id"
|
||||
|
|
|
|||
|
|
@ -59,6 +59,7 @@ class TestS3VectorsVectorStoreConfig:
|
|||
api_base="https://s3vectors.us-west-2.api.aws",
|
||||
litellm_logging_obj=mock_logging_obj,
|
||||
litellm_params={},
|
||||
extra_body=None,
|
||||
)
|
||||
|
||||
def test_transform_search_response(self):
|
||||
|
|
|
|||
|
|
@ -4261,3 +4261,32 @@ def test_sync_streaming_uses_custom_client():
|
|||
# Verify that gemini_client is in the partial's keywords
|
||||
assert "gemini_client" in partial_make_sync_call.keywords
|
||||
assert partial_make_sync_call.keywords["gemini_client"] is mock_client
|
||||
|
||||
|
||||
def test_transform_response_does_not_leak_body_on_parse_failure():
|
||||
leaky_body = {"candidates": [{"content": {"parts": [{"text": "secret content"}]}}]}
|
||||
raw_response = MagicMock()
|
||||
raw_response.json.return_value = leaky_body
|
||||
raw_response.text = json.dumps(leaky_body)
|
||||
raw_response.headers = {}
|
||||
|
||||
with patch(
|
||||
"litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini.GenerateContentResponseBody",
|
||||
side_effect=KeyError("missing required field"),
|
||||
):
|
||||
with pytest.raises(VertexAIError) as exc_info:
|
||||
VertexGeminiConfig().transform_response(
|
||||
model="gemini-pro",
|
||||
raw_response=raw_response,
|
||||
model_response=ModelResponse(),
|
||||
logging_obj=MagicMock(),
|
||||
request_data={},
|
||||
messages=[],
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
encoding=None,
|
||||
)
|
||||
|
||||
msg = str(exc_info.value)
|
||||
assert "secret content" not in msg
|
||||
assert "Error converting to valid response block" in msg
|
||||
|
|
|
|||
|
|
@ -728,6 +728,136 @@ class TestMCPServerManager:
|
|||
]
|
||||
assert scopes == ["read", "write"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_descovery_metadata_probes_well_known_when_server_does_not_challenge(
|
||||
self,
|
||||
):
|
||||
manager = MCPServerManager()
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.get = AsyncMock(return_value=mock_response)
|
||||
|
||||
mock_metadata = MCPOAuthMetadata(
|
||||
scopes=None,
|
||||
authorization_url="https://login.microsoftonline.com/tenant/oauth2/v2.0/authorize",
|
||||
token_url="https://login.microsoftonline.com/tenant/oauth2/v2.0/token",
|
||||
registration_url=None,
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
|
||||
return_value=mock_client,
|
||||
),
|
||||
patch.object(
|
||||
manager,
|
||||
"_attempt_well_known_discovery",
|
||||
AsyncMock(
|
||||
return_value=(
|
||||
["https://login.microsoftonline.com/test-tenant-id/v2.0"],
|
||||
["api://some-scope/.default"],
|
||||
)
|
||||
),
|
||||
) as mock_well_known,
|
||||
patch.object(
|
||||
manager,
|
||||
"_fetch_authorization_server_metadata",
|
||||
AsyncMock(return_value=mock_metadata),
|
||||
) as mock_fetch_auth,
|
||||
):
|
||||
result = await manager._descovery_metadata("http://localhost:8001/mcp")
|
||||
|
||||
mock_well_known.assert_awaited_once_with("http://localhost:8001/mcp")
|
||||
mock_fetch_auth.assert_awaited_once_with(
|
||||
["https://login.microsoftonline.com/test-tenant-id/v2.0"]
|
||||
)
|
||||
assert result is mock_metadata
|
||||
assert result.scopes == ["api://some-scope/.default"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_single_authorization_server_metadata_supports_azure_issuer_path(
|
||||
self,
|
||||
):
|
||||
manager = MCPServerManager()
|
||||
issuer = "https://login.microsoftonline.com/test-tenant-id/v2.0"
|
||||
|
||||
def build_response(url: str):
|
||||
mock_response = MagicMock()
|
||||
if url == f"{issuer}/.well-known/openid-configuration":
|
||||
mock_response.json.return_value = {
|
||||
"authorization_endpoint": "https://login.microsoftonline.com/test-tenant-id/oauth2/v2.0/authorize",
|
||||
"token_endpoint": "https://login.microsoftonline.com/test-tenant-id/oauth2/v2.0/token",
|
||||
"scopes_supported": ["api://some-scope/.default"],
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
else:
|
||||
request = httpx.Request("GET", url)
|
||||
response_obj = httpx.Response(status_code=404, request=request)
|
||||
mock_response.raise_for_status = MagicMock(
|
||||
side_effect=httpx.HTTPStatusError(
|
||||
"not found", request=request, response=response_obj
|
||||
)
|
||||
)
|
||||
return mock_response
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.get = AsyncMock(side_effect=build_response)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
|
||||
return_value=mock_client,
|
||||
):
|
||||
result = await manager._fetch_single_authorization_server_metadata(issuer)
|
||||
|
||||
assert result is not None
|
||||
assert (
|
||||
result.authorization_url
|
||||
== "https://login.microsoftonline.com/test-tenant-id/oauth2/v2.0/authorize"
|
||||
)
|
||||
assert (
|
||||
result.token_url
|
||||
== "https://login.microsoftonline.com/test-tenant-id/oauth2/v2.0/token"
|
||||
)
|
||||
assert result.scopes == ["api://some-scope/.default"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fetch_single_authorization_server_metadata_derives_azure_metadata(
|
||||
self,
|
||||
):
|
||||
manager = MCPServerManager()
|
||||
issuer = "https://login.microsoftonline.com/test-tenant-id/v2.0"
|
||||
|
||||
request = httpx.Request("GET", issuer)
|
||||
response_obj = httpx.Response(status_code=404, request=request)
|
||||
mock_response = MagicMock()
|
||||
mock_response.raise_for_status = MagicMock(
|
||||
side_effect=httpx.HTTPStatusError(
|
||||
"not found", request=request, response=response_obj
|
||||
)
|
||||
)
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.get = AsyncMock(return_value=mock_response)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.get_async_httpx_client",
|
||||
return_value=mock_client,
|
||||
):
|
||||
result = await manager._fetch_single_authorization_server_metadata(issuer)
|
||||
|
||||
assert result is not None
|
||||
assert (
|
||||
result.authorization_url
|
||||
== "https://login.microsoftonline.com/test-tenant-id/oauth2/v2.0/authorize"
|
||||
)
|
||||
assert (
|
||||
result.token_url
|
||||
== "https://login.microsoftonline.com/test-tenant-id/oauth2/v2.0/token"
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_descovery_metadata_falls_back_to_origin_when_no_auth_servers(self):
|
||||
manager = MCPServerManager()
|
||||
|
|
|
|||
|
|
@ -0,0 +1,302 @@
|
|||
"""
|
||||
Tests for the short-ID MCP tool prefix (LITELLM_USE_SHORT_MCP_TOOL_PREFIX).
|
||||
|
||||
The short-prefix mode swaps the historical alias/server_name prefix on
|
||||
tool names for a deterministic three-character base62 ID derived from the
|
||||
server's ``server_id``. This keeps tool names well below the 60-char
|
||||
upper bound enforced by some model APIs while remaining stable across
|
||||
processes/restarts and tolerant of mixed-version clients.
|
||||
"""
|
||||
|
||||
from typing import List
|
||||
|
||||
import pytest
|
||||
from mcp.types import Tool as MCPTool
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
SHORT_MCP_TOOL_PREFIX_LENGTH,
|
||||
add_server_prefix_to_name,
|
||||
compute_short_server_prefix,
|
||||
get_server_prefix,
|
||||
is_short_mcp_tool_prefix_enabled,
|
||||
iter_known_server_prefixes,
|
||||
)
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
|
||||
def _make_server(
|
||||
*,
|
||||
server_id: str = "abcdef-1234",
|
||||
server_name: str = "github_onprem",
|
||||
alias: str = "github_onprem",
|
||||
) -> MCPServer:
|
||||
return MCPServer(
|
||||
server_id=server_id,
|
||||
name=alias or server_name,
|
||||
alias=alias,
|
||||
server_name=server_name,
|
||||
transport="http",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_env(monkeypatch):
|
||||
monkeypatch.delenv("LITELLM_USE_SHORT_MCP_TOOL_PREFIX", raising=False)
|
||||
yield
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Pure helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestShortPrefixHelpers:
|
||||
def test_short_prefix_is_three_base62_chars(self):
|
||||
prefix = compute_short_server_prefix("any-server-id")
|
||||
assert len(prefix) == SHORT_MCP_TOOL_PREFIX_LENGTH
|
||||
assert prefix.isalnum() and prefix.isascii()
|
||||
|
||||
def test_short_prefix_first_char_is_alphabetic(self):
|
||||
"""The first char must be [A-Za-z] so the prefix is a valid identifier
|
||||
on every model API (some providers historically required the first
|
||||
character of a function name to be alphabetic)."""
|
||||
# Sweep many server_ids and rehash attempts to give us coverage of
|
||||
# every position the high-order bits can land on.
|
||||
for i in range(200):
|
||||
for attempt in range(4):
|
||||
prefix = compute_short_server_prefix(f"server-{i}", attempt=attempt)
|
||||
assert prefix[0].isalpha(), (
|
||||
f"prefix {prefix!r} for server-{i} (attempt={attempt}) "
|
||||
f"starts with a non-alphabetic character"
|
||||
)
|
||||
|
||||
def test_short_prefix_is_deterministic(self):
|
||||
assert compute_short_server_prefix("abc") == compute_short_server_prefix("abc")
|
||||
assert compute_short_server_prefix("abc") != compute_short_server_prefix("abd")
|
||||
|
||||
def test_short_prefix_requires_server_id(self):
|
||||
with pytest.raises(ValueError):
|
||||
compute_short_server_prefix("")
|
||||
|
||||
def test_flag_defaults_to_false(self):
|
||||
assert is_short_mcp_tool_prefix_enabled() is False
|
||||
|
||||
@pytest.mark.parametrize("value", ["1", "true", "TRUE", "yes", "On"])
|
||||
def test_flag_truthy_values(self, monkeypatch, value):
|
||||
monkeypatch.setenv("LITELLM_USE_SHORT_MCP_TOOL_PREFIX", value)
|
||||
assert is_short_mcp_tool_prefix_enabled() is True
|
||||
|
||||
@pytest.mark.parametrize("value", ["0", "false", "no", "off", ""])
|
||||
def test_flag_falsey_values(self, monkeypatch, value):
|
||||
monkeypatch.setenv("LITELLM_USE_SHORT_MCP_TOOL_PREFIX", value)
|
||||
assert is_short_mcp_tool_prefix_enabled() is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_server_prefix behaviour
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGetServerPrefix:
|
||||
def test_default_mode_uses_alias(self):
|
||||
server = _make_server(alias="github_onprem", server_name="github_onprem")
|
||||
assert get_server_prefix(server) == "github_onprem"
|
||||
|
||||
def test_short_mode_uses_short_id(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_USE_SHORT_MCP_TOOL_PREFIX", "true")
|
||||
server = _make_server(server_id="abcdef-1234")
|
||||
prefix = get_server_prefix(server)
|
||||
assert prefix == compute_short_server_prefix("abcdef-1234")
|
||||
assert len(prefix) == SHORT_MCP_TOOL_PREFIX_LENGTH
|
||||
|
||||
def test_short_mode_falls_back_when_no_server_id(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_USE_SHORT_MCP_TOOL_PREFIX", "true")
|
||||
|
||||
class _Bare:
|
||||
alias = "fallback_alias"
|
||||
server_name = None
|
||||
server_id = None
|
||||
|
||||
assert get_server_prefix(_Bare()) == "fallback_alias"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# iter_known_server_prefixes — covers reverse-lookup tolerance
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestIterKnownServerPrefixes:
|
||||
def test_default_mode_includes_short_id_too(self):
|
||||
server = _make_server()
|
||||
prefixes = list(iter_known_server_prefixes(server))
|
||||
# Contains the live prefix and every known form so that mixed-mode
|
||||
# clients can be resolved.
|
||||
assert "github_onprem" in prefixes
|
||||
assert compute_short_server_prefix(server.server_id) in prefixes
|
||||
|
||||
def test_short_mode_still_yields_long_forms(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_USE_SHORT_MCP_TOOL_PREFIX", "true")
|
||||
server = _make_server()
|
||||
prefixes = list(iter_known_server_prefixes(server))
|
||||
assert "github_onprem" in prefixes
|
||||
assert compute_short_server_prefix(server.server_id) in prefixes
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Manager-level behaviour: list + reverse-lookup
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _stub_tools() -> List[MCPTool]:
|
||||
return [
|
||||
MCPTool(name="get_repo", description="", inputSchema={"type": "object"}),
|
||||
MCPTool(name="list_issues", description="", inputSchema={"type": "object"}),
|
||||
]
|
||||
|
||||
|
||||
class TestManagerShortPrefix:
|
||||
def test_list_tools_uses_short_prefix_when_flag_on(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_USE_SHORT_MCP_TOOL_PREFIX", "true")
|
||||
manager = MCPServerManager()
|
||||
server = _make_server()
|
||||
|
||||
out = manager._create_prefixed_tools(_stub_tools(), server)
|
||||
|
||||
short = compute_short_server_prefix(server.server_id)
|
||||
assert {t.name for t in out} == {f"{short}-get_repo", f"{short}-list_issues"}
|
||||
|
||||
def test_call_tool_lookup_resolves_short_prefix(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_USE_SHORT_MCP_TOOL_PREFIX", "true")
|
||||
manager = MCPServerManager()
|
||||
server = _make_server()
|
||||
manager.registry[server.server_id] = server
|
||||
manager._create_prefixed_tools(_stub_tools(), server)
|
||||
|
||||
short = compute_short_server_prefix(server.server_id)
|
||||
resolved = manager._get_mcp_server_from_tool_name(f"{short}-get_repo")
|
||||
assert resolved is server
|
||||
|
||||
def test_call_tool_lookup_resolves_long_prefix_in_short_mode(self, monkeypatch):
|
||||
"""Old clients that cached the long-prefix name must still route."""
|
||||
monkeypatch.setenv("LITELLM_USE_SHORT_MCP_TOOL_PREFIX", "true")
|
||||
manager = MCPServerManager()
|
||||
server = _make_server()
|
||||
manager.registry[server.server_id] = server
|
||||
manager._create_prefixed_tools(_stub_tools(), server)
|
||||
|
||||
resolved = manager._get_mcp_server_from_tool_name("github_onprem-get_repo")
|
||||
assert resolved is server
|
||||
|
||||
def test_default_mode_unchanged(self):
|
||||
manager = MCPServerManager()
|
||||
server = _make_server()
|
||||
|
||||
out = manager._create_prefixed_tools(_stub_tools(), server)
|
||||
|
||||
assert {t.name for t in out} == {
|
||||
"github_onprem-get_repo",
|
||||
"github_onprem-list_issues",
|
||||
}
|
||||
assert (
|
||||
manager._get_mcp_server_from_tool_name("github_onprem-get_repo") is None
|
||||
) # registry empty
|
||||
manager.registry[server.server_id] = server
|
||||
assert (
|
||||
manager._get_mcp_server_from_tool_name("github_onprem-get_repo") is server
|
||||
)
|
||||
|
||||
def test_total_tool_name_length_short_enough(self, monkeypatch):
|
||||
"""The short prefix keeps tool names under the 60-char limit even
|
||||
when the upstream tool name is itself reasonably long."""
|
||||
monkeypatch.setenv("LITELLM_USE_SHORT_MCP_TOOL_PREFIX", "true")
|
||||
long_server_name = "a" * 50
|
||||
server = _make_server(
|
||||
server_id="server-id-1",
|
||||
server_name=long_server_name,
|
||||
alias=long_server_name,
|
||||
)
|
||||
prefix = get_server_prefix(server)
|
||||
full = add_server_prefix_to_name("get_repo", prefix)
|
||||
assert len(full) < 60
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Collision-resolution at registration time
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestShortPrefixCollisionResolution:
|
||||
"""``_assign_unique_short_prefix`` must rehash on collision.
|
||||
|
||||
The dedup path is exercised by forcing two distinct ``server_id``
|
||||
values to both hash to the same natural prefix via a monkeypatched
|
||||
``compute_short_server_prefix``.
|
||||
"""
|
||||
|
||||
def test_no_op_when_flag_off(self):
|
||||
manager = MCPServerManager()
|
||||
server = _make_server(server_id="abc")
|
||||
manager._assign_unique_short_prefix(server)
|
||||
assert server.short_prefix is None
|
||||
|
||||
def test_assigns_natural_hash_when_no_collision(self, monkeypatch):
|
||||
from litellm.proxy._experimental.mcp_server import utils as mcp_utils
|
||||
|
||||
monkeypatch.setenv("LITELLM_USE_SHORT_MCP_TOOL_PREFIX", "true")
|
||||
manager = MCPServerManager()
|
||||
server = _make_server(server_id="abc")
|
||||
manager._assign_unique_short_prefix(server)
|
||||
|
||||
assert server.short_prefix == mcp_utils.compute_short_server_prefix("abc")
|
||||
|
||||
def test_rehashes_when_natural_hash_collides(self, monkeypatch):
|
||||
"""Two server_ids that natural-hash to the same prefix get
|
||||
deterministic, distinct short prefixes."""
|
||||
monkeypatch.setenv("LITELLM_USE_SHORT_MCP_TOOL_PREFIX", "true")
|
||||
|
||||
# Force every attempt=0 hash to "AAA" and attempt=1 to "AAB".
|
||||
# That way the second server registered must rehash to "AAB".
|
||||
from litellm.proxy._experimental.mcp_server import utils as mcp_utils
|
||||
|
||||
def _fake_hash(server_id: str, attempt: int = 0) -> str:
|
||||
return "AAA" if attempt == 0 else f"AA{chr(ord('A') + attempt)}"
|
||||
|
||||
monkeypatch.setattr(mcp_utils, "compute_short_server_prefix", _fake_hash)
|
||||
# Also patch the symbol that the manager imported at module load.
|
||||
from litellm.proxy._experimental.mcp_server import (
|
||||
mcp_server_manager as mgr_module,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(mgr_module, "compute_short_server_prefix", _fake_hash)
|
||||
|
||||
manager = MCPServerManager()
|
||||
first = _make_server(server_id="server-1", alias="srv1")
|
||||
second = _make_server(server_id="server-2", alias="srv2")
|
||||
|
||||
# Pretend both are already in the registry so dedup sees both.
|
||||
manager.registry[first.server_id] = first
|
||||
manager._assign_unique_short_prefix(first)
|
||||
manager.registry[second.server_id] = second
|
||||
manager._assign_unique_short_prefix(second)
|
||||
|
||||
assert first.short_prefix == "AAA"
|
||||
assert second.short_prefix == "AAB"
|
||||
assert first.short_prefix != second.short_prefix
|
||||
|
||||
def test_cached_prefix_is_reused(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_USE_SHORT_MCP_TOOL_PREFIX", "true")
|
||||
manager = MCPServerManager()
|
||||
server = _make_server(server_id="abc")
|
||||
server.short_prefix = "ZZZ" # pretend a previous registration set this
|
||||
|
||||
manager._assign_unique_short_prefix(server)
|
||||
|
||||
assert server.short_prefix == "ZZZ"
|
||||
|
||||
def test_get_server_prefix_prefers_cached(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_USE_SHORT_MCP_TOOL_PREFIX", "true")
|
||||
server = _make_server(server_id="abc")
|
||||
server.short_prefix = "Q9q"
|
||||
|
||||
assert get_server_prefix(server) == "Q9q"
|
||||
|
|
@ -0,0 +1,255 @@
|
|||
"""
|
||||
Unit tests for `call_with_db_reconnect_retry` — the canonical "try DB read,
|
||||
on transport error reconnect once and retry once" helper.
|
||||
|
||||
Covers the regression in issue #25143 where read paths (e.g.
|
||||
`PrismaClient.get_generic_data`) lost their reconnect-and-retry-once branch in
|
||||
LiteLLM 1.83.x and started emitting `db_exceptions` alerts on transient
|
||||
`httpx.ReadError` flaps that used to self-heal in 1.82.6.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from prisma.errors import UniqueViolationError
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry
|
||||
|
||||
|
||||
def _make_client(
|
||||
*,
|
||||
attempt_db_reconnect_return: bool = True,
|
||||
has_attempt_db_reconnect: bool = True,
|
||||
):
|
||||
"""Build a minimal stand-in for PrismaClient that exposes only the surface
|
||||
`call_with_db_reconnect_retry` actually pokes at."""
|
||||
client = MagicMock()
|
||||
if has_attempt_db_reconnect:
|
||||
client.attempt_db_reconnect = AsyncMock(
|
||||
return_value=attempt_db_reconnect_return
|
||||
)
|
||||
else:
|
||||
# `hasattr(client, "attempt_db_reconnect")` must return False — MagicMock
|
||||
# auto-creates attributes, so we wipe it out via `spec`.
|
||||
client = MagicMock(spec=[])
|
||||
client._db_auth_reconnect_timeout_seconds = 2.0
|
||||
client._db_auth_reconnect_lock_timeout_seconds = 0.1
|
||||
return client
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_with_db_reconnect_retry_returns_value_on_first_success():
|
||||
"""Happy path: factory succeeds first call, no reconnect attempted."""
|
||||
client = _make_client()
|
||||
|
||||
async def _factory():
|
||||
return {"id": 1}
|
||||
|
||||
result = await call_with_db_reconnect_retry(client, _factory, reason="happy_path")
|
||||
|
||||
assert result == {"id": 1}
|
||||
client.attempt_db_reconnect.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_with_db_reconnect_retry_retries_after_transport_error():
|
||||
"""Transport error on first call → reconnect → second call succeeds."""
|
||||
client = _make_client(attempt_db_reconnect_return=True)
|
||||
|
||||
invocations = []
|
||||
|
||||
async def _factory():
|
||||
invocations.append(None)
|
||||
if len(invocations) == 1:
|
||||
raise httpx.ReadError("transport blip")
|
||||
return {"id": 1}
|
||||
|
||||
result = await call_with_db_reconnect_retry(
|
||||
client, _factory, reason="prisma_get_generic_data_config_lookup_failure"
|
||||
)
|
||||
|
||||
assert result == {"id": 1}
|
||||
assert len(invocations) == 2
|
||||
client.attempt_db_reconnect.assert_awaited_once()
|
||||
call_kwargs = client.attempt_db_reconnect.await_args.kwargs
|
||||
assert call_kwargs["reason"] == "prisma_get_generic_data_config_lookup_failure"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_with_db_reconnect_retry_does_not_retry_on_data_layer_error():
|
||||
"""Data-layer errors (e.g. UniqueViolationError) are NOT transport errors —
|
||||
propagate immediately, do not reconnect."""
|
||||
client = _make_client()
|
||||
|
||||
async def _factory():
|
||||
raise UniqueViolationError(
|
||||
data={"user_facing_error": {"meta": {}}},
|
||||
message="Unique constraint failed",
|
||||
)
|
||||
|
||||
with pytest.raises(UniqueViolationError):
|
||||
await call_with_db_reconnect_retry(client, _factory, reason="data_layer_test")
|
||||
|
||||
client.attempt_db_reconnect.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_with_db_reconnect_retry_propagates_when_reconnect_fails():
|
||||
"""Transport error, but reconnect returns False → propagate the original
|
||||
exception. Do not call factory a second time."""
|
||||
client = _make_client(attempt_db_reconnect_return=False)
|
||||
|
||||
invocations = []
|
||||
|
||||
async def _factory():
|
||||
invocations.append(None)
|
||||
raise httpx.ReadError("transport blip")
|
||||
|
||||
with pytest.raises(httpx.ReadError):
|
||||
await call_with_db_reconnect_retry(client, _factory, reason="reconnect_fails")
|
||||
|
||||
assert len(invocations) == 1
|
||||
client.attempt_db_reconnect.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_with_db_reconnect_retry_propagates_after_second_transport_error():
|
||||
"""Transport error, reconnect succeeds, retry also raises transport error →
|
||||
propagate. At most one retry by construction (no infinite loop)."""
|
||||
client = _make_client(attempt_db_reconnect_return=True)
|
||||
|
||||
invocations = []
|
||||
|
||||
async def _factory():
|
||||
invocations.append(None)
|
||||
raise httpx.ReadError("still failing")
|
||||
|
||||
with pytest.raises(httpx.ReadError):
|
||||
await call_with_db_reconnect_retry(
|
||||
client, _factory, reason="second_transport_error"
|
||||
)
|
||||
|
||||
assert len(invocations) == 2
|
||||
client.attempt_db_reconnect.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_with_db_reconnect_retry_skips_when_no_attempt_db_reconnect_attr():
|
||||
"""Older PrismaClient stand-ins / partial mocks may not expose
|
||||
`attempt_db_reconnect`. The helper must not crash — just propagate the
|
||||
original exception. Mirrors the `hasattr` guard from
|
||||
`auth_checks._fetch_key_object_from_db_with_reconnect`."""
|
||||
client = _make_client(has_attempt_db_reconnect=False)
|
||||
|
||||
async def _factory():
|
||||
raise httpx.ReadError("transport blip")
|
||||
|
||||
with pytest.raises(httpx.ReadError):
|
||||
await call_with_db_reconnect_retry(client, _factory, reason="no_reconnect_attr")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_with_db_reconnect_retry_invokes_factory_twice_not_same_coro():
|
||||
"""Guard against the obvious bug of awaiting the same coroutine twice
|
||||
(`RuntimeError: cannot reuse already awaited coroutine`). The helper must
|
||||
call the factory a fresh time on retry, not cache an awaitable."""
|
||||
client = _make_client(attempt_db_reconnect_return=True)
|
||||
|
||||
factory_call_count = 0
|
||||
|
||||
async def _factory():
|
||||
nonlocal factory_call_count
|
||||
factory_call_count += 1
|
||||
if factory_call_count == 1:
|
||||
raise httpx.ReadError("transport blip")
|
||||
return "ok"
|
||||
|
||||
result = await call_with_db_reconnect_retry(
|
||||
client, _factory, reason="fresh_coro_on_retry"
|
||||
)
|
||||
|
||||
assert result == "ok"
|
||||
assert factory_call_count == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_with_db_reconnect_retry_passes_explicit_timeouts():
|
||||
"""Explicit timeout_seconds / lock_timeout_seconds override the auth
|
||||
defaults read off the prisma_client object."""
|
||||
client = _make_client(attempt_db_reconnect_return=True)
|
||||
|
||||
async def _factory():
|
||||
if not hasattr(_factory, "_called"):
|
||||
_factory._called = True # type: ignore[attr-defined]
|
||||
raise httpx.ReadError("transport blip")
|
||||
return "ok"
|
||||
|
||||
result = await call_with_db_reconnect_retry(
|
||||
client,
|
||||
_factory,
|
||||
reason="explicit_timeouts",
|
||||
timeout_seconds=5.5,
|
||||
lock_timeout_seconds=0.25,
|
||||
)
|
||||
|
||||
assert result == "ok"
|
||||
call_kwargs = client.attempt_db_reconnect.await_args.kwargs
|
||||
assert call_kwargs["timeout_seconds"] == 5.5
|
||||
assert call_kwargs["lock_timeout_seconds"] == 0.25
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_with_db_reconnect_retry_uses_auth_defaults_when_unset():
|
||||
"""When timeouts are not provided, helper reads
|
||||
`_db_auth_reconnect_timeout_seconds` / `_db_auth_reconnect_lock_timeout_seconds`
|
||||
off the prisma_client (matching the auth path's existing convention)."""
|
||||
client = _make_client(attempt_db_reconnect_return=True)
|
||||
client._db_auth_reconnect_timeout_seconds = 3.0
|
||||
client._db_auth_reconnect_lock_timeout_seconds = 0.5
|
||||
|
||||
async def _factory():
|
||||
if not hasattr(_factory, "_called"):
|
||||
_factory._called = True # type: ignore[attr-defined]
|
||||
raise httpx.ReadError("transport blip")
|
||||
return "ok"
|
||||
|
||||
await call_with_db_reconnect_retry(client, _factory, reason="defaults")
|
||||
|
||||
call_kwargs = client.attempt_db_reconnect.await_args.kwargs
|
||||
assert call_kwargs["timeout_seconds"] == 3.0
|
||||
assert call_kwargs["lock_timeout_seconds"] == 0.5
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_with_db_reconnect_retry_preserves_original_error_when_reconnect_raises():
|
||||
"""If `attempt_db_reconnect` itself raises (lock cancellation, timer
|
||||
error, unexpected internal failure), the helper must surface the
|
||||
*original* transport error to telemetry — not the reconnect exception.
|
||||
Otherwise `failure_handler` / `db_exceptions` alerts log the wrong
|
||||
error string and the actual DB transport problem becomes invisible.
|
||||
|
||||
The reconnect error is chained as the `__cause__` for debuggability."""
|
||||
client = MagicMock()
|
||||
reconnect_exc = RuntimeError("simulated reconnect lock cancellation")
|
||||
client.attempt_db_reconnect = AsyncMock(side_effect=reconnect_exc)
|
||||
client._db_auth_reconnect_timeout_seconds = 2.0
|
||||
client._db_auth_reconnect_lock_timeout_seconds = 0.1
|
||||
|
||||
original_exc = httpx.ReadError("transport blip")
|
||||
|
||||
async def _factory():
|
||||
raise original_exc
|
||||
|
||||
with pytest.raises(httpx.ReadError) as exc_info:
|
||||
await call_with_db_reconnect_retry(
|
||||
client, _factory, reason="reconnect_itself_raises"
|
||||
)
|
||||
|
||||
assert exc_info.value is original_exc
|
||||
assert exc_info.value.__cause__ is reconnect_exc
|
||||
client.attempt_db_reconnect.assert_awaited_once()
|
||||
|
|
@ -5,6 +5,7 @@ import sys
|
|||
import time
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
sys.path.insert(
|
||||
|
|
@ -34,18 +35,18 @@ async def test_attempt_db_reconnect_should_succeed(mock_proxy_logging):
|
|||
client = PrismaClient(
|
||||
database_url="mock://test", proxy_logging_obj=mock_proxy_logging
|
||||
)
|
||||
client.db.disconnect = AsyncMock(return_value=None)
|
||||
client.db.connect = AsyncMock(return_value=None)
|
||||
client.db.recreate_prisma_client = AsyncMock(return_value=None)
|
||||
client.db.query_raw = AsyncMock(return_value=[{"result": 1}])
|
||||
client._start_engine_watcher = AsyncMock()
|
||||
|
||||
result = await client.attempt_db_reconnect(
|
||||
reason="unit_test_reconnect_success",
|
||||
force=True,
|
||||
)
|
||||
with patch.dict(os.environ, {"DATABASE_URL": "postgresql://test"}):
|
||||
result = await client.attempt_db_reconnect(
|
||||
reason="unit_test_reconnect_success",
|
||||
force=True,
|
||||
)
|
||||
|
||||
assert result is True
|
||||
client.db.disconnect.assert_awaited_once()
|
||||
client.db.connect.assert_awaited_once()
|
||||
client.db.recreate_prisma_client.assert_awaited_once_with("postgresql://test")
|
||||
client.db.query_raw.assert_awaited_once_with("SELECT 1")
|
||||
|
||||
|
||||
|
|
@ -140,15 +141,19 @@ async def test_attempt_db_reconnect_should_set_cooldown_after_attempt(
|
|||
)
|
||||
client._db_last_reconnect_attempt_ts = 0.0
|
||||
client._db_reconnect_cooldown_seconds = 10
|
||||
client.db.disconnect = AsyncMock(return_value=None)
|
||||
client.db.connect = AsyncMock(return_value=None)
|
||||
client.db.recreate_prisma_client = AsyncMock(return_value=None)
|
||||
client.db.query_raw = AsyncMock(return_value=[{"result": 1}])
|
||||
client._start_engine_watcher = AsyncMock()
|
||||
|
||||
# Use a counter-based mock to avoid StopIteration when time.time() is called
|
||||
# more times than expected (varies by Python version / internal code paths).
|
||||
fake_clock = iter(range(100, 10000))
|
||||
with patch(
|
||||
"litellm.proxy.utils.time.time", side_effect=lambda: float(next(fake_clock))
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.utils.time.time",
|
||||
side_effect=lambda: float(next(fake_clock)),
|
||||
),
|
||||
patch.dict(os.environ, {"DATABASE_URL": "postgresql://test"}),
|
||||
):
|
||||
result = await client.attempt_db_reconnect(
|
||||
reason="unit_test_cooldown_timestamp_after_attempt",
|
||||
|
|
@ -162,23 +167,28 @@ async def test_attempt_db_reconnect_should_set_cooldown_after_attempt(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_reconnect_cycle_watchdog_should_use_direct_db_ops(
|
||||
async def test_run_reconnect_cycle_watchdog_should_use_recreate_prisma_client(
|
||||
mock_proxy_logging,
|
||||
):
|
||||
"""Direct reconnect goes through recreate_prisma_client (which non-blockingly
|
||||
kills the old engine) instead of calling disconnect() — see issue #26191.
|
||||
"""
|
||||
client = PrismaClient(
|
||||
database_url="mock://test", proxy_logging_obj=mock_proxy_logging
|
||||
)
|
||||
client.disconnect = AsyncMock(side_effect=AssertionError("wrapper disconnect used"))
|
||||
client.connect = AsyncMock(side_effect=AssertionError("wrapper connect used"))
|
||||
client.db.disconnect = AsyncMock(return_value=None)
|
||||
client.db.connect = AsyncMock(return_value=None)
|
||||
client.db.disconnect = AsyncMock(
|
||||
side_effect=AssertionError("disconnect must not be called")
|
||||
)
|
||||
client.db.recreate_prisma_client = AsyncMock(return_value=None)
|
||||
client.db.query_raw = AsyncMock(return_value=[{"result": 1}])
|
||||
client._start_engine_watcher = AsyncMock()
|
||||
|
||||
await client._run_reconnect_cycle(timeout_seconds=None)
|
||||
with patch.dict(os.environ, {"DATABASE_URL": "postgresql://test"}):
|
||||
await client._run_reconnect_cycle(timeout_seconds=None)
|
||||
|
||||
client.db.disconnect.assert_awaited_once()
|
||||
client.db.connect.assert_awaited_once()
|
||||
client.db.recreate_prisma_client.assert_awaited_once_with("postgresql://test")
|
||||
client.db.query_raw.assert_awaited_once_with("SELECT 1")
|
||||
client.db.disconnect.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -189,19 +199,22 @@ async def test_run_reconnect_cycle_watchdog_should_use_default_timeout_budget(
|
|||
database_url="mock://test", proxy_logging_obj=mock_proxy_logging
|
||||
)
|
||||
client._db_watchdog_reconnect_timeout_seconds = 0.1
|
||||
client.db.disconnect = AsyncMock(return_value=None)
|
||||
client._start_engine_watcher = AsyncMock()
|
||||
|
||||
async def _slow_connect():
|
||||
async def _slow_recreate(_db_url):
|
||||
await asyncio.sleep(0.08)
|
||||
|
||||
async def _slow_query(_query: str):
|
||||
await asyncio.sleep(0.08)
|
||||
return [{"result": 1}]
|
||||
|
||||
client.db.connect = AsyncMock(side_effect=_slow_connect)
|
||||
client.db.recreate_prisma_client = AsyncMock(side_effect=_slow_recreate)
|
||||
client.db.query_raw = AsyncMock(side_effect=_slow_query)
|
||||
|
||||
with pytest.raises(asyncio.TimeoutError):
|
||||
with (
|
||||
pytest.raises(asyncio.TimeoutError),
|
||||
patch.dict(os.environ, {"DATABASE_URL": "postgresql://test"}),
|
||||
):
|
||||
await client._run_reconnect_cycle(timeout_seconds=None)
|
||||
|
||||
|
||||
|
|
@ -212,19 +225,22 @@ async def test_run_reconnect_cycle_timeout_should_use_single_overall_budget(
|
|||
client = PrismaClient(
|
||||
database_url="mock://test", proxy_logging_obj=mock_proxy_logging
|
||||
)
|
||||
client.db.disconnect = AsyncMock(return_value=None)
|
||||
client._start_engine_watcher = AsyncMock()
|
||||
|
||||
async def _slow_connect():
|
||||
async def _slow_recreate(_db_url):
|
||||
await asyncio.sleep(0.08)
|
||||
|
||||
async def _slow_query(_query: str):
|
||||
await asyncio.sleep(0.08)
|
||||
return [{"result": 1}]
|
||||
|
||||
client.db.connect = AsyncMock(side_effect=_slow_connect)
|
||||
client.db.recreate_prisma_client = AsyncMock(side_effect=_slow_recreate)
|
||||
client.db.query_raw = AsyncMock(side_effect=_slow_query)
|
||||
|
||||
with pytest.raises(asyncio.TimeoutError):
|
||||
with (
|
||||
pytest.raises(asyncio.TimeoutError),
|
||||
patch.dict(os.environ, {"DATABASE_URL": "postgresql://test"}),
|
||||
):
|
||||
await client._run_reconnect_cycle(timeout_seconds=0.1)
|
||||
|
||||
|
||||
|
|
@ -319,42 +335,154 @@ async def test_db_health_watchdog_start_stop_lifecycle(mock_proxy_logging):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lightweight_reconnect_kills_engine_on_disconnect_failure(
|
||||
async def test_recreate_prisma_client_kills_old_engine_without_disconnect(
|
||||
mock_proxy_logging,
|
||||
):
|
||||
"""Lightweight reconnect must kill the old engine PID when disconnect() fails."""
|
||||
"""recreate_prisma_client SIGTERMs the old engine PID directly rather than
|
||||
calling `disconnect()`, which blocks the asyncio event loop on the sync
|
||||
`subprocess.Popen.wait()` inside prisma-client-py — see issue #26191.
|
||||
"""
|
||||
client = PrismaClient(
|
||||
database_url="mock://test", proxy_logging_obj=mock_proxy_logging
|
||||
)
|
||||
client.db.disconnect = AsyncMock(side_effect=Exception("disconnect failed"))
|
||||
client.db.connect = AsyncMock(return_value=None)
|
||||
client.db.query_raw = AsyncMock(return_value=[{"result": 1}])
|
||||
disconnect_mock = AsyncMock(
|
||||
side_effect=AssertionError("disconnect must not be called on reconnect path")
|
||||
)
|
||||
client.db._original_prisma.disconnect = disconnect_mock
|
||||
|
||||
with (
|
||||
patch.object(client, "_get_engine_pid", return_value=9999),
|
||||
patch("os.kill") as mock_kill,
|
||||
patch("asyncio.sleep", new_callable=AsyncMock),
|
||||
patch.object(client.db, "_get_engine_pid", return_value=9999),
|
||||
patch("litellm.proxy.db.prisma_client.os.kill") as mock_kill,
|
||||
patch("litellm.proxy.db.prisma_client.asyncio.sleep", new_callable=AsyncMock),
|
||||
):
|
||||
await client._run_reconnect_cycle(timeout_seconds=5.0)
|
||||
# Return a Prisma instance whose connect() is awaitable.
|
||||
fake_new_prisma = MagicMock()
|
||||
fake_new_prisma.connect = AsyncMock(return_value=None)
|
||||
with patch("prisma.Prisma", return_value=fake_new_prisma):
|
||||
await client.db.recreate_prisma_client("postgresql://test")
|
||||
|
||||
mock_kill.assert_any_call(9999, signal.SIGTERM)
|
||||
client.db.connect.assert_awaited_once()
|
||||
client.db.query_raw.assert_awaited_once_with("SELECT 1")
|
||||
disconnect_mock.assert_not_awaited()
|
||||
fake_new_prisma.connect.assert_awaited_once()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# get_generic_data: transport-reconnect-and-retry coverage (issue #25143)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lightweight_reconnect_skips_kill_on_successful_disconnect(
|
||||
async def test_get_generic_data_retries_on_transport_error_for_config_table(
|
||||
mock_proxy_logging,
|
||||
):
|
||||
"""Lightweight reconnect must NOT kill when disconnect() succeeds."""
|
||||
"""`get_generic_data(table_name="config")` self-heals on a transient
|
||||
`httpx.ReadError`: reconnect once, retry once, return the row.
|
||||
|
||||
Regression for issue #25143 — the 1.83.x line lost the reconnect-and-retry
|
||||
branch that 1.82.6 had on this method. `_update_config_from_db` fans out
|
||||
four concurrent `get_generic_data` calls, so a single transport flap used
|
||||
to surface as four `db_exceptions` alerts and a stale config window.
|
||||
"""
|
||||
client = PrismaClient(
|
||||
database_url="mock://test", proxy_logging_obj=mock_proxy_logging
|
||||
)
|
||||
client.db.disconnect = AsyncMock(return_value=None)
|
||||
client.db.connect = AsyncMock(return_value=None)
|
||||
client.db.query_raw = AsyncMock(return_value=[{"result": 1}])
|
||||
|
||||
with patch("os.kill") as mock_kill:
|
||||
await client._run_reconnect_cycle(timeout_seconds=5.0)
|
||||
expected_row = {"param_name": "general_settings", "param_value": {"foo": "bar"}}
|
||||
invocations: list[None] = []
|
||||
|
||||
mock_kill.assert_not_called()
|
||||
async def _flaky_find_first(**kwargs):
|
||||
invocations.append(None)
|
||||
if len(invocations) == 1:
|
||||
raise httpx.ReadError("simulated transport blip")
|
||||
return expected_row
|
||||
|
||||
client.db.litellm_config.find_first = AsyncMock(side_effect=_flaky_find_first)
|
||||
client.attempt_db_reconnect = AsyncMock(return_value=True)
|
||||
|
||||
result = await client.get_generic_data(
|
||||
key="param_name",
|
||||
value="general_settings",
|
||||
table_name="config",
|
||||
)
|
||||
|
||||
assert result == expected_row
|
||||
assert len(invocations) == 2
|
||||
client.attempt_db_reconnect.assert_awaited_once()
|
||||
reconnect_kwargs = client.attempt_db_reconnect.await_args.kwargs
|
||||
assert reconnect_kwargs["reason"] == "prisma_get_generic_data_config_lookup_failure"
|
||||
|
||||
# The failure_handler telemetry side-effect must NOT fire on the first
|
||||
# transport blip — only if the post-retry call also fails. Drain the
|
||||
# event loop so any spuriously-spawned task would have run by now.
|
||||
await asyncio.sleep(0)
|
||||
mock_proxy_logging.failure_handler.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_generic_data_propagates_when_reconnect_fails(mock_proxy_logging):
|
||||
"""If reconnect itself does not succeed, propagate the original transport
|
||||
error and let the existing failure_handler / db_exceptions telemetry fire."""
|
||||
client = PrismaClient(
|
||||
database_url="mock://test", proxy_logging_obj=mock_proxy_logging
|
||||
)
|
||||
|
||||
client.db.litellm_config.find_first = AsyncMock(
|
||||
side_effect=httpx.ReadError("simulated transport blip")
|
||||
)
|
||||
client.attempt_db_reconnect = AsyncMock(return_value=False)
|
||||
|
||||
with pytest.raises(httpx.ReadError):
|
||||
await client.get_generic_data(
|
||||
key="param_name",
|
||||
value="general_settings",
|
||||
table_name="config",
|
||||
)
|
||||
|
||||
client.attempt_db_reconnect.assert_awaited_once()
|
||||
# Failure telemetry IS expected here — the read genuinely failed.
|
||||
await asyncio.sleep(0)
|
||||
mock_proxy_logging.failure_handler.assert_called_once()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _engine_confirmed_dead flag-reset bug (B2)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_engine_confirmed_dead_persists_across_failed_heavy_reconnect(
|
||||
mock_proxy_logging,
|
||||
):
|
||||
"""Regression test for the flag-reset bug.
|
||||
|
||||
Before the fix, `_run_reconnect_cycle` cleared
|
||||
`self._engine_confirmed_dead = False` *before* awaiting
|
||||
`_do_heavy_reconnect()`. If the heavy reconnect raised (e.g. timeout,
|
||||
missing DATABASE_URL, recreate failure), the flag was left cleared and the
|
||||
next attempt could demote to the lightweight path even though the engine
|
||||
was genuinely dead.
|
||||
|
||||
The fix moves the reset into the success branch — the flag must stay True
|
||||
when heavy reconnect raises.
|
||||
"""
|
||||
client = PrismaClient(
|
||||
database_url="mock://test", proxy_logging_obj=mock_proxy_logging
|
||||
)
|
||||
client._engine_confirmed_dead = True
|
||||
client._engine_pid = 0 # so `_is_engine_alive` is not consulted
|
||||
|
||||
# Make the heavy reconnect path raise.
|
||||
client.db.recreate_prisma_client = AsyncMock(
|
||||
side_effect=RuntimeError("simulated heavy reconnect failure")
|
||||
)
|
||||
client._start_engine_watcher = AsyncMock()
|
||||
client._cleanup_engine_watcher = MagicMock()
|
||||
client._reap_all_zombies = MagicMock()
|
||||
|
||||
with patch.dict(os.environ, {"DATABASE_URL": "postgresql://test"}):
|
||||
with pytest.raises(Exception):
|
||||
await client._run_reconnect_cycle(timeout_seconds=5.0)
|
||||
|
||||
# The flag must STILL be True so the next attempt re-enters the heavy
|
||||
# branch instead of silently demoting to the lightweight path.
|
||||
assert client._engine_confirmed_dead is True
|
||||
|
|
|
|||
|
|
@ -0,0 +1,611 @@
|
|||
"""
|
||||
Unit tests for workflow management endpoints (/v1/workflows/runs/*).
|
||||
Uses FastAPI TestClient with a mocked prisma_client.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
from prisma.errors import UniqueViolationError
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
from litellm.proxy.management_endpoints.workflow_management_endpoints import router
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _make_run(
|
||||
run_id: str = "run-1",
|
||||
session_id: str = "sess-1",
|
||||
workflow_type: str = "shin-builder",
|
||||
status: str = "pending",
|
||||
created_by: Any = "tok-test",
|
||||
) -> MagicMock:
|
||||
obj = MagicMock()
|
||||
obj.run_id = run_id
|
||||
obj.session_id = session_id
|
||||
obj.workflow_type = workflow_type
|
||||
obj.status = status
|
||||
obj.created_by = created_by
|
||||
obj.created_at = datetime.now(timezone.utc)
|
||||
obj.updated_at = datetime.now(timezone.utc)
|
||||
obj.input = None
|
||||
obj.output = None
|
||||
obj.metadata = None
|
||||
return obj
|
||||
|
||||
|
||||
def _make_event(
|
||||
event_id: str = "evt-1",
|
||||
run_id: str = "run-1",
|
||||
event_type: str = "step.started",
|
||||
step_name: str = "grill",
|
||||
sequence_number: int = 0,
|
||||
) -> MagicMock:
|
||||
obj = MagicMock()
|
||||
obj.event_id = event_id
|
||||
obj.run_id = run_id
|
||||
obj.event_type = event_type
|
||||
obj.step_name = step_name
|
||||
obj.sequence_number = sequence_number
|
||||
obj.data = None
|
||||
obj.created_at = datetime.now(timezone.utc)
|
||||
return obj
|
||||
|
||||
|
||||
def _make_message(
|
||||
message_id: str = "msg-1",
|
||||
run_id: str = "run-1",
|
||||
role: str = "user",
|
||||
content: str = "hello",
|
||||
sequence_number: int = 0,
|
||||
) -> MagicMock:
|
||||
obj = MagicMock()
|
||||
obj.message_id = message_id
|
||||
obj.run_id = run_id
|
||||
obj.role = role
|
||||
obj.content = content
|
||||
obj.sequence_number = sequence_number
|
||||
obj.session_id = None
|
||||
obj.created_at = datetime.now(timezone.utc)
|
||||
return obj
|
||||
|
||||
|
||||
def _make_tx(event_return=None, run_return=None, msg_return=None) -> MagicMock:
|
||||
"""Build an async context-manager mock for prisma_client.db.tx()."""
|
||||
tx = MagicMock()
|
||||
tx.litellm_workflowevent = MagicMock()
|
||||
tx.litellm_workflowevent.create = AsyncMock(
|
||||
return_value=event_return or _make_event()
|
||||
)
|
||||
tx.litellm_workflowrun = MagicMock()
|
||||
tx.litellm_workflowrun.update = AsyncMock(return_value=run_return or _make_run())
|
||||
tx.litellm_workflowmessage = MagicMock()
|
||||
tx.litellm_workflowmessage.create = AsyncMock(
|
||||
return_value=msg_return or _make_message()
|
||||
)
|
||||
tx.__aenter__ = AsyncMock(return_value=tx)
|
||||
tx.__aexit__ = AsyncMock(return_value=False)
|
||||
return tx
|
||||
|
||||
|
||||
def _make_prisma_client() -> MagicMock:
|
||||
client = MagicMock()
|
||||
client.db = MagicMock()
|
||||
client.db.litellm_workflowrun = MagicMock()
|
||||
client.db.litellm_workflowevent = MagicMock()
|
||||
client.db.litellm_workflowmessage = MagicMock()
|
||||
# default tx() returns a no-op transaction
|
||||
client.db.tx = MagicMock(return_value=_make_tx())
|
||||
return client
|
||||
|
||||
|
||||
def _make_app() -> FastAPI:
|
||||
app = FastAPI()
|
||||
app.include_router(router)
|
||||
return app
|
||||
|
||||
|
||||
def _override_auth() -> Any:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
auth = UserAPIKeyAuth(api_key="sk-test", user_id="admin")
|
||||
auth.token = "tok-test"
|
||||
return auth
|
||||
|
||||
|
||||
def _override_auth_admin() -> Any:
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
|
||||
auth = UserAPIKeyAuth(api_key="sk-master")
|
||||
auth.user_role = LitellmUserRoles.PROXY_ADMIN # type: ignore[assignment]
|
||||
return auth
|
||||
|
||||
|
||||
def _override_auth_user_with_token(token: str = "tok-abc") -> Any:
|
||||
"""Return a non-admin caller whose hashed token equals `token`."""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
auth = UserAPIKeyAuth(api_key="sk-user", user_id="user-1")
|
||||
auth.token = token # override the computed hash with a predictable value
|
||||
return auth
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCreateWorkflowRun:
|
||||
def setup_method(self):
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
self._prisma = _make_prisma_client()
|
||||
app = _make_app()
|
||||
app.dependency_overrides[user_api_key_auth] = _override_auth
|
||||
self.client = TestClient(app, raise_server_exceptions=True)
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_create_returns_run(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.create = AsyncMock(return_value=_make_run())
|
||||
|
||||
resp = self.client.post(
|
||||
"/v1/workflows/runs",
|
||||
json={"workflow_type": "shin-builder"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
self._prisma.db.litellm_workflowrun.create.assert_awaited_once()
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client", None)
|
||||
def test_create_500_when_no_db(self):
|
||||
resp = self.client.post(
|
||||
"/v1/workflows/runs",
|
||||
json={"workflow_type": "shin-builder"},
|
||||
)
|
||||
assert resp.status_code == 500
|
||||
|
||||
|
||||
class TestListWorkflowRuns:
|
||||
def setup_method(self):
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
self._prisma = _make_prisma_client()
|
||||
app = _make_app()
|
||||
app.dependency_overrides[user_api_key_auth] = _override_auth
|
||||
self.client = TestClient(app, raise_server_exceptions=True)
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_list_returns_runs(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_many = AsyncMock(
|
||||
return_value=[_make_run()]
|
||||
)
|
||||
|
||||
resp = self.client.get("/v1/workflows/runs")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["count"] == 1
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_list_filters_by_status(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_many = AsyncMock(return_value=[])
|
||||
|
||||
resp = self.client.get("/v1/workflows/runs?status=running")
|
||||
assert resp.status_code == 200
|
||||
call_kwargs = self._prisma.db.litellm_workflowrun.find_many.call_args[1]
|
||||
assert call_kwargs["where"]["status"] == "running"
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_list_filters_by_multiple_statuses(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_many = AsyncMock(return_value=[])
|
||||
|
||||
resp = self.client.get("/v1/workflows/runs?status=running,paused")
|
||||
assert resp.status_code == 200
|
||||
call_kwargs = self._prisma.db.litellm_workflowrun.find_many.call_args[1]
|
||||
assert call_kwargs["where"]["status"] == {"in": ["running", "paused"]}
|
||||
|
||||
|
||||
class TestGetWorkflowRun:
|
||||
def setup_method(self):
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
self._prisma = _make_prisma_client()
|
||||
app = _make_app()
|
||||
app.dependency_overrides[user_api_key_auth] = _override_auth
|
||||
self.client = TestClient(app, raise_server_exceptions=True)
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_get_existing_run(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(
|
||||
return_value=_make_run()
|
||||
)
|
||||
|
||||
resp = self.client.get("/v1/workflows/runs/run-1")
|
||||
assert resp.status_code == 200
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_get_missing_run_returns_404(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
resp = self.client.get("/v1/workflows/runs/nonexistent")
|
||||
assert resp.status_code == 404
|
||||
|
||||
|
||||
class TestUpdateWorkflowRun:
|
||||
def setup_method(self):
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
self._prisma = _make_prisma_client()
|
||||
app = _make_app()
|
||||
app.dependency_overrides[user_api_key_auth] = _override_auth
|
||||
self.client = TestClient(app, raise_server_exceptions=True)
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_update_status(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(
|
||||
return_value=_make_run()
|
||||
)
|
||||
updated = _make_run(status="completed")
|
||||
self._prisma.db.litellm_workflowrun.update = AsyncMock(return_value=updated)
|
||||
|
||||
resp = self.client.patch(
|
||||
"/v1/workflows/runs/run-1", json={"status": "completed"}
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
self._prisma.db.litellm_workflowrun.update.assert_awaited_once()
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_update_no_fields_returns_400(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
resp = self.client.patch("/v1/workflows/runs/run-1", json={})
|
||||
assert resp.status_code == 400
|
||||
|
||||
|
||||
class TestAppendWorkflowEvent:
|
||||
def setup_method(self):
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
self._prisma = _make_prisma_client()
|
||||
app = _make_app()
|
||||
app.dependency_overrides[user_api_key_auth] = _override_auth
|
||||
self.client = TestClient(app, raise_server_exceptions=True)
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_append_event_updates_run_status(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
# _require_run check
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(
|
||||
return_value=_make_run()
|
||||
)
|
||||
self._prisma.db.litellm_workflowevent.find_many = AsyncMock(return_value=[])
|
||||
tx = _make_tx(
|
||||
event_return=_make_event(), run_return=_make_run(status="running")
|
||||
)
|
||||
self._prisma.db.tx = MagicMock(return_value=tx)
|
||||
|
||||
resp = self.client.post(
|
||||
"/v1/workflows/runs/run-1/events",
|
||||
json={"event_type": "step.started", "step_name": "grill"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
# run status updated inside tx
|
||||
tx.litellm_workflowrun.update.assert_awaited_once()
|
||||
update_call = tx.litellm_workflowrun.update.call_args[1]
|
||||
assert update_call["data"]["status"] == "running"
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_append_event_no_status_update_for_unknown_type(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(
|
||||
return_value=_make_run()
|
||||
)
|
||||
self._prisma.db.litellm_workflowevent.find_many = AsyncMock(return_value=[])
|
||||
tx = _make_tx(event_return=_make_event(event_type="custom.event"))
|
||||
self._prisma.db.tx = MagicMock(return_value=tx)
|
||||
|
||||
resp = self.client.post(
|
||||
"/v1/workflows/runs/run-1/events",
|
||||
json={"event_type": "custom.event", "step_name": "grill"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
# no status update inside tx for unknown event_type
|
||||
tx.litellm_workflowrun.update.assert_not_awaited()
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_sequence_number_increments(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(
|
||||
return_value=_make_run()
|
||||
)
|
||||
existing = _make_event(sequence_number=4)
|
||||
self._prisma.db.litellm_workflowevent.find_many = AsyncMock(
|
||||
return_value=[existing]
|
||||
)
|
||||
tx = _make_tx(event_return=_make_event(sequence_number=5))
|
||||
self._prisma.db.tx = MagicMock(return_value=tx)
|
||||
|
||||
self.client.post(
|
||||
"/v1/workflows/runs/run-1/events",
|
||||
json={"event_type": "step.started", "step_name": "plan"},
|
||||
)
|
||||
create_call = tx.litellm_workflowevent.create.call_args[1]
|
||||
assert create_call["data"]["sequence_number"] == 5
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_unknown_run_id_returns_404(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
resp = self.client.post(
|
||||
"/v1/workflows/runs/nonexistent/events",
|
||||
json={"event_type": "step.started", "step_name": "grill"},
|
||||
)
|
||||
assert resp.status_code == 404
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_sequence_collision_retries_and_succeeds(self, mock_pc):
|
||||
"""UniqueViolationError on first attempt triggers retry; second attempt succeeds."""
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(
|
||||
return_value=_make_run()
|
||||
)
|
||||
self._prisma.db.litellm_workflowevent.find_many = AsyncMock(return_value=[])
|
||||
|
||||
# First tx raises UniqueViolationError; second succeeds.
|
||||
tx_fail = _make_tx()
|
||||
tx_fail.__aenter__ = AsyncMock(return_value=tx_fail)
|
||||
tx_fail.litellm_workflowevent.create = AsyncMock(
|
||||
side_effect=UniqueViolationError(
|
||||
{"user_facing_error": {"message": "unique"}}
|
||||
)
|
||||
)
|
||||
tx_fail.__aexit__ = AsyncMock(return_value=False)
|
||||
|
||||
tx_ok = _make_tx(event_return=_make_event(sequence_number=1))
|
||||
|
||||
self._prisma.db.tx = MagicMock(side_effect=[tx_fail, tx_ok])
|
||||
|
||||
resp = self.client.post(
|
||||
"/v1/workflows/runs/run-1/events",
|
||||
json={"event_type": "step.started", "step_name": "grill"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
|
||||
|
||||
class TestWorkflowMessages:
|
||||
def setup_method(self):
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
self._prisma = _make_prisma_client()
|
||||
app = _make_app()
|
||||
app.dependency_overrides[user_api_key_auth] = _override_auth
|
||||
self.client = TestClient(app, raise_server_exceptions=True)
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_append_message(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(
|
||||
return_value=_make_run()
|
||||
)
|
||||
self._prisma.db.litellm_workflowmessage.find_many = AsyncMock(return_value=[])
|
||||
self._prisma.db.litellm_workflowmessage.create = AsyncMock(
|
||||
return_value=_make_message()
|
||||
)
|
||||
|
||||
resp = self.client.post(
|
||||
"/v1/workflows/runs/run-1/messages",
|
||||
json={"role": "user", "content": "fix the bug"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_append_message_unknown_run_returns_404(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
resp = self.client.post(
|
||||
"/v1/workflows/runs/nonexistent/messages",
|
||||
json={"role": "user", "content": "hello"},
|
||||
)
|
||||
assert resp.status_code == 404
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_list_messages_ordered(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(
|
||||
return_value=_make_run()
|
||||
)
|
||||
self._prisma.db.litellm_workflowmessage.find_many = AsyncMock(
|
||||
return_value=[
|
||||
_make_message(sequence_number=0),
|
||||
_make_message(sequence_number=1, role="assistant"),
|
||||
]
|
||||
)
|
||||
|
||||
resp = self.client.get("/v1/workflows/runs/run-1/messages")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["count"] == 2
|
||||
call_kwargs = self._prisma.db.litellm_workflowmessage.find_many.call_args[1]
|
||||
assert call_kwargs["order"] == {"sequence_number": "asc"}
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_list_messages_respects_limit(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(
|
||||
return_value=_make_run()
|
||||
)
|
||||
self._prisma.db.litellm_workflowmessage.find_many = AsyncMock(return_value=[])
|
||||
|
||||
resp = self.client.get("/v1/workflows/runs/run-1/messages?limit=25")
|
||||
assert resp.status_code == 200
|
||||
call_kwargs = self._prisma.db.litellm_workflowmessage.find_many.call_args[1]
|
||||
assert call_kwargs["take"] == 25
|
||||
|
||||
|
||||
class TestListWorkflowEvents:
|
||||
def setup_method(self):
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
self._prisma = _make_prisma_client()
|
||||
app = _make_app()
|
||||
app.dependency_overrides[user_api_key_auth] = _override_auth
|
||||
self.client = TestClient(app, raise_server_exceptions=True)
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_list_events_ordered(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(
|
||||
return_value=_make_run()
|
||||
)
|
||||
self._prisma.db.litellm_workflowevent.find_many = AsyncMock(
|
||||
return_value=[
|
||||
_make_event(sequence_number=0),
|
||||
_make_event(sequence_number=1),
|
||||
]
|
||||
)
|
||||
|
||||
resp = self.client.get("/v1/workflows/runs/run-1/events")
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["count"] == 2
|
||||
call_kwargs = self._prisma.db.litellm_workflowevent.find_many.call_args[1]
|
||||
assert call_kwargs["order"] == {"sequence_number": "asc"}
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_list_events_respects_limit(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(
|
||||
return_value=_make_run()
|
||||
)
|
||||
self._prisma.db.litellm_workflowevent.find_many = AsyncMock(return_value=[])
|
||||
|
||||
resp = self.client.get("/v1/workflows/runs/run-1/events?limit=10")
|
||||
assert resp.status_code == 200
|
||||
call_kwargs = self._prisma.db.litellm_workflowevent.find_many.call_args[1]
|
||||
assert call_kwargs["take"] == 10
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_list_events_unknown_run_returns_404(self, mock_pc):
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
resp = self.client.get("/v1/workflows/runs/nonexistent/events")
|
||||
assert resp.status_code == 404
|
||||
|
||||
|
||||
class TestTenantIsolation:
|
||||
"""Ownership enforcement: non-admin callers only see their own runs."""
|
||||
|
||||
def _make_app_with_auth(self, auth_fn):
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
self._prisma = _make_prisma_client()
|
||||
app = _make_app()
|
||||
app.dependency_overrides[user_api_key_auth] = auth_fn
|
||||
return TestClient(app, raise_server_exceptions=True)
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_create_stores_caller_token(self, mock_pc):
|
||||
token = "tok-owner"
|
||||
client = self._make_app_with_auth(lambda: _override_auth_user_with_token(token))
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.create = AsyncMock(
|
||||
return_value=_make_run(created_by=token)
|
||||
)
|
||||
|
||||
resp = client.post("/v1/workflows/runs", json={"workflow_type": "test"})
|
||||
assert resp.status_code == 200
|
||||
create_call = self._prisma.db.litellm_workflowrun.create.call_args[1]
|
||||
assert create_call["data"]["created_by"] == token
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_non_admin_list_scoped_to_caller_token(self, mock_pc):
|
||||
token = "tok-owner"
|
||||
client = self._make_app_with_auth(lambda: _override_auth_user_with_token(token))
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_many = AsyncMock(return_value=[])
|
||||
|
||||
resp = client.get("/v1/workflows/runs")
|
||||
assert resp.status_code == 200
|
||||
call_kwargs = self._prisma.db.litellm_workflowrun.find_many.call_args[1]
|
||||
assert call_kwargs["where"].get("created_by") == token
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_admin_list_not_scoped(self, mock_pc):
|
||||
client = self._make_app_with_auth(_override_auth_admin)
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_many = AsyncMock(return_value=[])
|
||||
|
||||
resp = client.get("/v1/workflows/runs")
|
||||
assert resp.status_code == 200
|
||||
call_kwargs = self._prisma.db.litellm_workflowrun.find_many.call_args[1]
|
||||
assert "created_by" not in call_kwargs["where"]
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_non_admin_get_other_users_run_returns_404(self, mock_pc):
|
||||
token = "tok-caller"
|
||||
client = self._make_app_with_auth(lambda: _override_auth_user_with_token(token))
|
||||
mock_pc.db = self._prisma.db
|
||||
# Run owned by a different key
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(
|
||||
return_value=_make_run(created_by="tok-other-owner")
|
||||
)
|
||||
|
||||
resp = client.get("/v1/workflows/runs/run-1")
|
||||
assert resp.status_code == 404
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_non_admin_get_null_owner_run_returns_404(self, mock_pc):
|
||||
token = "tok-caller"
|
||||
client = self._make_app_with_auth(lambda: _override_auth_user_with_token(token))
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(
|
||||
return_value=_make_run(created_by=None)
|
||||
)
|
||||
|
||||
resp = client.get("/v1/workflows/runs/run-1")
|
||||
assert resp.status_code == 404
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_non_admin_update_null_owner_run_returns_404(self, mock_pc):
|
||||
token = "tok-caller"
|
||||
client = self._make_app_with_auth(lambda: _override_auth_user_with_token(token))
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(
|
||||
return_value=_make_run(created_by=None)
|
||||
)
|
||||
self._prisma.db.litellm_workflowrun.update = AsyncMock(
|
||||
return_value=_make_run(status="completed")
|
||||
)
|
||||
|
||||
resp = client.patch("/v1/workflows/runs/run-1", json={"status": "completed"})
|
||||
assert resp.status_code == 404
|
||||
self._prisma.db.litellm_workflowrun.update.assert_not_awaited()
|
||||
|
||||
@patch("litellm.proxy.proxy_server.prisma_client")
|
||||
def test_non_admin_get_own_run_succeeds(self, mock_pc):
|
||||
token = "tok-caller"
|
||||
client = self._make_app_with_auth(lambda: _override_auth_user_with_token(token))
|
||||
mock_pc.db = self._prisma.db
|
||||
self._prisma.db.litellm_workflowrun.find_unique = AsyncMock(
|
||||
return_value=_make_run(created_by=token)
|
||||
)
|
||||
|
||||
resp = client.get("/v1/workflows/runs/run-1")
|
||||
assert resp.status_code == 200
|
||||
|
|
@ -218,6 +218,141 @@ class TestProxyBaseLLMRequestProcessing:
|
|||
headers_with_invalid
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_litellm_proxy_success_headers_from_llm_response(self):
|
||||
"""
|
||||
Google native :generateContent uses this helper instead of base_process_llm_request;
|
||||
ensure x-litellm-* headers and callback hooks merge like the main proxy path.
|
||||
"""
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.headers = {}
|
||||
|
||||
class _FakeGenaiResponse:
|
||||
_hidden_params = {
|
||||
"model_id": "deployment-model-id",
|
||||
"cache_key": "ck-test",
|
||||
"api_base": "https://generativelanguage.googleapis.com/v1beta",
|
||||
"response_cost": 0.001,
|
||||
"additional_headers": {"llm_provider-ratelimit-requests": "1000"},
|
||||
}
|
||||
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.litellm_call_id = "call-id-test"
|
||||
|
||||
mock_user = MagicMock()
|
||||
mock_user.tpm_limit = None
|
||||
mock_user.rpm_limit = None
|
||||
mock_user.max_budget = None
|
||||
mock_user.spend = 0.0
|
||||
mock_user.allowed_model_region = None
|
||||
|
||||
proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||||
proxy_logging_obj.post_call_response_headers_hook = AsyncMock(
|
||||
return_value={"x-ratelimit-remaining-requests": "999"}
|
||||
)
|
||||
|
||||
headers = await ProxyBaseLLMRequestProcessing.build_litellm_proxy_success_headers_from_llm_response(
|
||||
response=_FakeGenaiResponse(),
|
||||
request_data={"model": "gemini/gemini-1.5-flash"},
|
||||
request=mock_request,
|
||||
user_api_key_dict=mock_user,
|
||||
logging_obj=logging_obj,
|
||||
version="9.9.9",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
assert headers["x-litellm-call-id"] == "call-id-test"
|
||||
assert headers["x-litellm-model-id"] == "deployment-model-id"
|
||||
assert headers["x-litellm-version"] == "9.9.9"
|
||||
assert headers["llm_provider-ratelimit-requests"] == "1000"
|
||||
assert headers["x-ratelimit-remaining-requests"] == "999"
|
||||
proxy_logging_obj.post_call_response_headers_hook.assert_awaited_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_litellm_proxy_success_headers_streaming_style_iterator(self):
|
||||
"""AsyncGoogleGenAIGenerateContentStreamingIterator sets _hidden_params at init; headers must propagate."""
|
||||
|
||||
class _FakeStreamLike:
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self):
|
||||
raise StopAsyncIteration
|
||||
|
||||
_hidden_params = {
|
||||
"model_id": "stream-model-id",
|
||||
"api_base": "https://generativelanguage.googleapis.com/v1beta",
|
||||
"cache_key": "",
|
||||
"response_cost": "",
|
||||
"additional_headers": {"llm_provider-x": "y"},
|
||||
}
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.headers = {}
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.litellm_call_id = "cid-stream"
|
||||
mock_user = MagicMock()
|
||||
mock_user.tpm_limit = None
|
||||
mock_user.rpm_limit = None
|
||||
mock_user.max_budget = None
|
||||
mock_user.spend = 0.0
|
||||
mock_user.allowed_model_region = None
|
||||
proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||||
proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value={})
|
||||
|
||||
headers = await ProxyBaseLLMRequestProcessing.build_litellm_proxy_success_headers_from_llm_response(
|
||||
response=_FakeStreamLike(),
|
||||
request_data={"model": "gemini/gemini-2.0-flash"},
|
||||
request=mock_request,
|
||||
user_api_key_dict=mock_user,
|
||||
logging_obj=logging_obj,
|
||||
version="1.0.0",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
assert headers["x-litellm-model-id"] == "stream-model-id"
|
||||
assert headers["x-litellm-model-api-base"] == (
|
||||
"https://generativelanguage.googleapis.com/v1beta"
|
||||
)
|
||||
assert headers["llm_provider-x"] == "y"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_litellm_proxy_success_headers_no_hidden_params_metadata_fallback(
|
||||
self,
|
||||
):
|
||||
"""When response has no _hidden_params, model_id can still come from litellm_metadata."""
|
||||
|
||||
class _BareResponse:
|
||||
pass
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.headers = {}
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.litellm_call_id = "cid-meta"
|
||||
mock_user = MagicMock()
|
||||
mock_user.tpm_limit = None
|
||||
mock_user.rpm_limit = None
|
||||
mock_user.max_budget = None
|
||||
mock_user.spend = 0.0
|
||||
mock_user.allowed_model_region = None
|
||||
proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||||
proxy_logging_obj.post_call_response_headers_hook = AsyncMock(return_value={})
|
||||
|
||||
headers = await ProxyBaseLLMRequestProcessing.build_litellm_proxy_success_headers_from_llm_response(
|
||||
response=_BareResponse(),
|
||||
request_data={
|
||||
"model": "gemini/gemini-1.5-flash",
|
||||
"litellm_metadata": {"model_info": {"id": "meta-model-id"}},
|
||||
},
|
||||
request=mock_request,
|
||||
user_api_key_dict=mock_user,
|
||||
logging_obj=logging_obj,
|
||||
version="1.0.0",
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
|
||||
assert headers["x-litellm-model-id"] == "meta-model-id"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_litellm_data_to_request_with_stream_timeout_header(self):
|
||||
"""
|
||||
|
|
@ -1158,13 +1293,6 @@ class TestCommonRequestProcessingHelpers:
|
|||
assert mock_tracer.trace.call_count == 4
|
||||
|
||||
# Verify that each call was made with the correct operation name
|
||||
expected_calls = [
|
||||
(("streaming.chunk.yield",), {}),
|
||||
(("streaming.chunk.yield",), {}),
|
||||
(("streaming.chunk.yield",), {}),
|
||||
(("streaming.chunk.yield",), {}),
|
||||
]
|
||||
|
||||
actual_calls = mock_tracer.trace.call_args_list
|
||||
assert len(actual_calls) == 4
|
||||
|
||||
|
|
|
|||
|
|
@ -29,8 +29,21 @@ from litellm.types.utils import (
|
|||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _use_local_model_cost_map(monkeypatch):
|
||||
original_model_cost = litellm.model_cost
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
litellm.get_model_info.cache_clear()
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
litellm.model_cost = original_model_cost
|
||||
litellm.get_model_info.cache_clear()
|
||||
|
||||
|
||||
class TestGPTImageCostCalculator:
|
||||
"""Test the OpenAI gpt-image-1 cost calculator"""
|
||||
"""Test the OpenAI gpt-image cost calculator"""
|
||||
|
||||
def test_gpt_image_1_cost_with_text_only(self):
|
||||
"""Test cost calculation with only text input tokens"""
|
||||
|
|
@ -149,6 +162,44 @@ class TestGPTImageCostCalculator:
|
|||
|
||||
assert cost == 0.0
|
||||
|
||||
def test_gpt_image_2_cost_with_text_and_image_tokens(self):
|
||||
"""Test cost calculation for gpt-image-2 token pricing"""
|
||||
from litellm.llms.openai.image_generation.cost_calculator import cost_calculator
|
||||
|
||||
usage = Usage(
|
||||
prompt_tokens=600,
|
||||
completion_tokens=5000,
|
||||
total_tokens=5600,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(
|
||||
text_tokens=100,
|
||||
image_tokens=500,
|
||||
),
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(
|
||||
text_tokens=1000,
|
||||
image_tokens=4000,
|
||||
),
|
||||
)
|
||||
|
||||
image_response = ImageResponse(
|
||||
created=1234567890,
|
||||
data=[ImageObject(url="http://example.com/image.jpg")],
|
||||
)
|
||||
image_response.usage = usage
|
||||
|
||||
cost = cost_calculator(
|
||||
model="gpt-image-2",
|
||||
image_response=image_response,
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
# GPT Image 2 pricing:
|
||||
# Text input: 100 * $5/1M = 0.0005
|
||||
# Image input: 500 * $8/1M = 0.004
|
||||
# Text output: 1000 * $10/1M = 0.01
|
||||
# Image output: 4000 * $30/1M = 0.12
|
||||
expected_cost = 0.0005 + 0.004 + 0.01 + 0.12
|
||||
assert abs(cost - expected_cost) < 1e-6, f"Expected {expected_cost}, got {cost}"
|
||||
|
||||
|
||||
class TestGPTImageCostRouting:
|
||||
"""Test that gpt-image models are properly routed to the token-based calculator"""
|
||||
|
|
@ -182,6 +233,33 @@ class TestGPTImageCostRouting:
|
|||
expected_cost = 0.0005 + 0.2
|
||||
assert abs(cost - expected_cost) < 1e-6, f"Expected {expected_cost}, got {cost}"
|
||||
|
||||
def test_openai_gpt_image_2_routes_to_token_calculator(self):
|
||||
"""Test that OpenAI gpt-image-2 routes to token-based calculator"""
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import CostCalculatorUtils
|
||||
|
||||
usage = Usage(
|
||||
prompt_tokens=100,
|
||||
completion_tokens=5000,
|
||||
total_tokens=5100,
|
||||
prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=100),
|
||||
completion_tokens_details=CompletionTokensDetailsWrapper(image_tokens=5000),
|
||||
)
|
||||
|
||||
image_response = ImageResponse(
|
||||
created=1234567890,
|
||||
data=[ImageObject(url="http://example.com/image.jpg")],
|
||||
)
|
||||
image_response.usage = usage
|
||||
|
||||
cost = CostCalculatorUtils.route_image_generation_cost_calculator(
|
||||
model="gpt-image-2",
|
||||
completion_response=image_response,
|
||||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
expected_cost = 0.0005 + 0.15
|
||||
assert abs(cost - expected_cost) < 1e-6, f"Expected {expected_cost}, got {cost}"
|
||||
|
||||
def test_openai_dalle_routes_to_pixel_calculator(self):
|
||||
"""Test that OpenAI DALL-E still routes to pixel-based calculator"""
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import CostCalculatorUtils
|
||||
|
|
|
|||
|
|
@ -32,6 +32,19 @@ from litellm.utils import (
|
|||
# Adds the parent directory to the system path
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def local_model_cost_map(monkeypatch):
|
||||
original_model_cost = litellm.model_cost
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
litellm.get_model_info.cache_clear()
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
litellm.model_cost = original_model_cost
|
||||
litellm.get_model_info.cache_clear()
|
||||
|
||||
|
||||
def test_check_provider_match_azure_ai_allows_openai_and_azure():
|
||||
"""
|
||||
Test that azure_ai provider can match openai and azure models.
|
||||
|
|
@ -198,6 +211,72 @@ def test_get_optional_params_image_gen_filters_empty_values():
|
|||
assert optional_params == {}
|
||||
|
||||
|
||||
def test_gpt_image_provider_detection_covers_existing_family():
|
||||
for image_model in ("gpt-image-1", "gpt-image-1-mini", "gpt-image-1.5"):
|
||||
model, custom_llm_provider, _, _ = litellm.get_llm_provider(model=image_model)
|
||||
|
||||
assert model == image_model
|
||||
assert custom_llm_provider == "openai"
|
||||
|
||||
|
||||
def test_gpt_image_2_provider_and_model_info(local_model_cost_map):
|
||||
|
||||
model, custom_llm_provider, _, _ = litellm.get_llm_provider(model="gpt-image-2")
|
||||
|
||||
assert model == "gpt-image-2"
|
||||
assert custom_llm_provider == "openai"
|
||||
|
||||
model_info = litellm.get_model_info(model="gpt-image-2")
|
||||
assert model_info["litellm_provider"] == "openai"
|
||||
assert model_info["mode"] == "image_generation"
|
||||
assert model_info["input_cost_per_token"] == 5e-06
|
||||
assert model_info["input_cost_per_image_token"] == 8e-06
|
||||
assert model_info["output_cost_per_token"] == 1e-05
|
||||
assert model_info["output_cost_per_image_token"] == 3e-05
|
||||
assert (
|
||||
"/v1/images/generations"
|
||||
in litellm.model_cost["gpt-image-2"]["supported_endpoints"]
|
||||
)
|
||||
assert (
|
||||
"/v1/images/edits" in litellm.model_cost["gpt-image-2"]["supported_endpoints"]
|
||||
)
|
||||
assert model_info["supports_vision"] is True
|
||||
assert model_info["supports_pdf_input"] is True
|
||||
|
||||
|
||||
def test_gpt_image_2_snapshot_model_info(local_model_cost_map):
|
||||
model, custom_llm_provider, _, _ = litellm.get_llm_provider(
|
||||
model="gpt-image-2-2026-04-21"
|
||||
)
|
||||
|
||||
assert model == "gpt-image-2-2026-04-21"
|
||||
assert custom_llm_provider == "openai"
|
||||
|
||||
model_info = litellm.get_model_info(model="gpt-image-2-2026-04-21")
|
||||
assert model_info["litellm_provider"] == "openai"
|
||||
assert model_info["mode"] == "image_generation"
|
||||
assert model_info["output_cost_per_image_token"] == 3e-05
|
||||
|
||||
|
||||
def test_azure_gpt_image_2_model_info(local_model_cost_map):
|
||||
model, custom_llm_provider, _, _ = litellm.get_llm_provider(
|
||||
model="azure/gpt-image-2"
|
||||
)
|
||||
|
||||
assert model == "gpt-image-2"
|
||||
assert custom_llm_provider == "azure"
|
||||
|
||||
model_info = litellm.get_model_info(
|
||||
model="gpt-image-2", custom_llm_provider="azure"
|
||||
)
|
||||
assert model_info["litellm_provider"] == "azure"
|
||||
assert model_info["mode"] == "image_generation"
|
||||
assert model_info["input_cost_per_token"] == 5e-06
|
||||
assert model_info["input_cost_per_image_token"] == 8e-06
|
||||
assert model_info["output_cost_per_token"] == 1e-05
|
||||
assert model_info["output_cost_per_image_token"] == 3e-05
|
||||
|
||||
|
||||
def test_all_model_configs():
|
||||
from litellm.llms.vertex_ai.vertex_ai_partner_models.ai21.transformation import (
|
||||
VertexAIAi21Config,
|
||||
|
|
|
|||
|
|
@ -267,6 +267,7 @@ class TestRAGFlowVectorStore(BaseVectorStoreTest):
|
|||
api_base="http://localhost:9380",
|
||||
litellm_logging_obj=logging_obj,
|
||||
litellm_params={},
|
||||
extra_body=None,
|
||||
)
|
||||
|
||||
def test_transform_search_vector_store_response_not_implemented(self):
|
||||
|
|
|
|||
|
|
@ -40,6 +40,7 @@ import { ProjectsPage } from "@/components/Projects/ProjectsPage";
|
|||
import VectorStoreManagement from "@/components/vector_store_management";
|
||||
import ToolPoliciesView from "@/components/ToolPoliciesView";
|
||||
import { MemoryView } from "@/components/MemoryView";
|
||||
import WorkflowRuns from "@/components/workflow_runs";
|
||||
import SpendLogsTable from "@/components/view_logs";
|
||||
import ViewUserDashboard from "@/components/view_users";
|
||||
import { ThemeProvider } from "@/contexts/ThemeContext";
|
||||
|
|
@ -631,6 +632,8 @@ function CreateKeyPageContent() {
|
|||
<VectorStoreManagement accessToken={accessToken} userRole={userRole} userID={userID} />
|
||||
) : page == "tool-policies" ? (
|
||||
<ToolPoliciesView accessToken={accessToken} userRole={userRole} />
|
||||
) : page == "workflows" ? (
|
||||
<WorkflowRuns accessToken={accessToken} />
|
||||
) : page == "memory" ? (
|
||||
<MemoryView
|
||||
accessToken={accessToken}
|
||||
|
|
|
|||
|
|
@ -69,7 +69,7 @@ describe("Sidebar (leftnav)", () => {
|
|||
"Virtual Keys",
|
||||
"Playground",
|
||||
"Models + Endpoints",
|
||||
"Agents",
|
||||
"Agentic",
|
||||
"MCP Servers",
|
||||
"Guardrails",
|
||||
"Policies",
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ import { useTeams } from "@/app/(dashboard)/hooks/teams/useTeams";
|
|||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import {
|
||||
ApiOutlined,
|
||||
ApartmentOutlined,
|
||||
AppstoreOutlined,
|
||||
AuditOutlined,
|
||||
BankOutlined,
|
||||
|
|
@ -120,11 +121,31 @@ const menuGroups: MenuGroup[] = [
|
|||
roles: rolesWithWriteAccess,
|
||||
},
|
||||
{
|
||||
key: "agents",
|
||||
page: "agents",
|
||||
label: "Agents",
|
||||
key: "agentic",
|
||||
page: "agentic",
|
||||
label: "Agentic",
|
||||
icon: <RobotOutlined />,
|
||||
roles: rolesWithWriteAccess,
|
||||
children: [
|
||||
{
|
||||
key: "agents",
|
||||
page: "agents",
|
||||
label: "Agents",
|
||||
icon: <RobotOutlined />,
|
||||
roles: rolesWithWriteAccess,
|
||||
},
|
||||
{
|
||||
key: "workflows",
|
||||
page: "workflows",
|
||||
label: "Workflow Runs",
|
||||
icon: <ApartmentOutlined />,
|
||||
},
|
||||
{
|
||||
key: "memory",
|
||||
page: "memory",
|
||||
label: "Memory",
|
||||
icon: <BookOutlined />,
|
||||
},
|
||||
],
|
||||
},
|
||||
{
|
||||
key: "mcp-servers",
|
||||
|
|
@ -139,12 +160,6 @@ const menuGroups: MenuGroup[] = [
|
|||
icon: <ApiOutlined />,
|
||||
roles: all_admin_roles,
|
||||
},
|
||||
{
|
||||
key: "memory",
|
||||
page: "memory",
|
||||
label: "Memory",
|
||||
icon: <BookOutlined />,
|
||||
},
|
||||
{
|
||||
key: "guardrails",
|
||||
page: "guardrails",
|
||||
|
|
|
|||
|
|
@ -9,6 +9,8 @@ export const pageDescriptions: Record<string, string> = {
|
|||
"llm-playground": "Interactive playground for testing LLM requests",
|
||||
models: "Configure and manage LLM models and endpoints",
|
||||
agents: "Create and manage AI agents",
|
||||
agentic: "Manage agentic resources: agents, workflow runs, and memory",
|
||||
workflows: "Track and inspect durable workflow run history",
|
||||
"mcp-servers": "Configure Model Context Protocol servers",
|
||||
memory: "Inspect and manage agent memory entries stored under /v1/memory",
|
||||
guardrails: "Set up content moderation and safety guardrails",
|
||||
|
|
|
|||
751
ui/litellm-dashboard/src/components/workflow_runs/index.tsx
Normal file
751
ui/litellm-dashboard/src/components/workflow_runs/index.tsx
Normal file
|
|
@ -0,0 +1,751 @@
|
|||
import React, { useState, useEffect, useCallback } from "react";
|
||||
import { Button, Collapse, Drawer, Empty, Spin, Table, Tooltip, Typography } from "antd";
|
||||
import { ReloadOutlined } from "@ant-design/icons";
|
||||
import { proxyBaseUrl } from "@/components/networking";
|
||||
|
||||
const { Text } = Typography;
|
||||
|
||||
interface WorkflowRunsProps {
|
||||
accessToken: string | null;
|
||||
}
|
||||
|
||||
type RunStatus = "pending" | "running" | "paused" | "completed" | "failed";
|
||||
|
||||
interface RunMetadata {
|
||||
title?: string;
|
||||
state?: string;
|
||||
pr_url?: string;
|
||||
worktree_path?: string;
|
||||
plan_text?: string;
|
||||
grill_session_id?: string;
|
||||
session_id?: string;
|
||||
[key: string]: unknown;
|
||||
}
|
||||
|
||||
interface WorkflowRun {
|
||||
run_id: string;
|
||||
status: RunStatus;
|
||||
workflow_type: string;
|
||||
created_at: string;
|
||||
metadata?: RunMetadata | null;
|
||||
}
|
||||
|
||||
interface WorkflowRunEvent {
|
||||
event_id: string;
|
||||
event_type: string;
|
||||
step_name: string;
|
||||
sequence_number: number;
|
||||
created_at: string;
|
||||
data?: Record<string, unknown> | null;
|
||||
}
|
||||
|
||||
interface WorkflowRunMessage {
|
||||
message_id: string;
|
||||
role: string;
|
||||
content: string;
|
||||
sequence_number: number;
|
||||
created_at: string;
|
||||
}
|
||||
|
||||
// ── design tokens ─────────────────────────────────────────────────────────────
|
||||
|
||||
const STATUS_DOT: Record<RunStatus, string> = {
|
||||
pending: "#a1a1aa",
|
||||
running: "#3b82f6",
|
||||
paused: "#f59e0b",
|
||||
completed: "#22c55e",
|
||||
failed: "#ef4444",
|
||||
};
|
||||
|
||||
const EVENT_COLOR: Record<string, { bar: string; border: string; text: string }> = {
|
||||
"step.started": { bar: "#f0fdf4", border: "#86efac", text: "#16a34a" },
|
||||
"step.failed": { bar: "#fef2f2", border: "#fca5a5", text: "#dc2626" },
|
||||
"hook.waiting": { bar: "#fffbeb", border: "#fcd34d", text: "#d97706" },
|
||||
"hook.received": { bar: "#eff6ff", border: "#93c5fd", text: "#2563eb" },
|
||||
};
|
||||
|
||||
function eventStyle(type: string) {
|
||||
return EVENT_COLOR[type] ?? { bar: "#f4f4f5", border: "#d4d4d8", text: "#52525b" };
|
||||
}
|
||||
|
||||
// ── helpers ───────────────────────────────────────────────────────────────────
|
||||
|
||||
function timeAgo(iso: string): string {
|
||||
const diff = Date.now() - new Date(iso).getTime();
|
||||
if (isNaN(diff)) return iso;
|
||||
const s = Math.floor(diff / 1000);
|
||||
if (s < 60) return `${s}s ago`;
|
||||
const m = Math.floor(s / 60);
|
||||
if (m < 60) return `${m}m ago`;
|
||||
const h = Math.floor(m / 60);
|
||||
if (h < 24) return `${h}h ago`;
|
||||
return `${Math.floor(h / 24)}d ago`;
|
||||
}
|
||||
|
||||
function fmtDuration(ms: number): string {
|
||||
if (ms < 0) return "";
|
||||
if (ms < 1000) return `${ms}ms`;
|
||||
return `${(ms / 1000).toFixed(1)}s`;
|
||||
}
|
||||
|
||||
function runTitle(run: WorkflowRun): string {
|
||||
const t = run.metadata?.title;
|
||||
if (t) return String(t);
|
||||
return run.workflow_type ?? run.run_id.slice(0, 8);
|
||||
}
|
||||
|
||||
function shortId(id: string): string {
|
||||
return id.slice(0, 8);
|
||||
}
|
||||
|
||||
// ── status dot ────────────────────────────────────────────────────────────────
|
||||
|
||||
const StatusDot: React.FC<{ status: RunStatus; size?: number }> = ({ status, size = 8 }) => (
|
||||
<span
|
||||
style={{
|
||||
display: "inline-block",
|
||||
width: size,
|
||||
height: size,
|
||||
borderRadius: "50%",
|
||||
background: STATUS_DOT[status] ?? "#a1a1aa",
|
||||
flexShrink: 0,
|
||||
}}
|
||||
/>
|
||||
);
|
||||
|
||||
// ── truncated text value ──────────────────────────────────────────────────────
|
||||
|
||||
const TRUNCATE_AT = 120;
|
||||
|
||||
const TruncatedValue: React.FC<{ value: string }> = ({ value }) => {
|
||||
const [expanded, setExpanded] = useState(false);
|
||||
if (value.length <= TRUNCATE_AT) {
|
||||
return <span style={{ color: "#27272a", wordBreak: "break-all" }}>{value}</span>;
|
||||
}
|
||||
return (
|
||||
<span style={{ color: "#27272a", wordBreak: "break-all" }}>
|
||||
{expanded ? value : value.slice(0, TRUNCATE_AT) + "…"}
|
||||
<button
|
||||
onClick={() => setExpanded((e) => !e)}
|
||||
style={{
|
||||
background: "none",
|
||||
border: "none",
|
||||
padding: "0 4px",
|
||||
cursor: "pointer",
|
||||
color: "#2563eb",
|
||||
fontSize: 11,
|
||||
flexShrink: 0,
|
||||
}}
|
||||
>
|
||||
{expanded ? "less" : "more"}
|
||||
</button>
|
||||
</span>
|
||||
);
|
||||
};
|
||||
|
||||
// ── metadata card ─────────────────────────────────────────────────────────────
|
||||
|
||||
const MetadataCard: React.FC<{ run: WorkflowRun }> = ({ run }) => {
|
||||
const meta = run.metadata ?? {};
|
||||
|
||||
const primaryFields: { key: string; label: string }[] = [
|
||||
{ key: "state", label: "state" },
|
||||
{ key: "worktree_path", label: "worktree" },
|
||||
{ key: "grill_session_id", label: "grill session" },
|
||||
{ key: "session_id", label: "session" },
|
||||
];
|
||||
|
||||
const primaryKeys = new Set(["title", ...primaryFields.map((f) => f.key)]);
|
||||
const extraEntries = Object.entries(meta).filter(
|
||||
([k, v]) => !primaryKeys.has(k) && v !== null && v !== undefined && v !== ""
|
||||
);
|
||||
|
||||
return (
|
||||
<div
|
||||
style={{
|
||||
borderRadius: 8,
|
||||
border: "1px solid #e4e4e7",
|
||||
marginBottom: 16,
|
||||
overflow: "hidden",
|
||||
}}
|
||||
>
|
||||
{/* title bar */}
|
||||
<div
|
||||
style={{
|
||||
padding: "14px 20px",
|
||||
borderBottom: "1px solid #f4f4f5",
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
gap: 10,
|
||||
}}
|
||||
>
|
||||
<StatusDot status={run.status} size={10} />
|
||||
<span style={{ fontSize: 14, fontWeight: 600, color: "#18181b", flex: 1 }}>
|
||||
{runTitle(run)}
|
||||
</span>
|
||||
<span
|
||||
style={{
|
||||
fontFamily: "monospace",
|
||||
fontSize: 11,
|
||||
color: "#a1a1aa",
|
||||
background: "#f4f4f5",
|
||||
padding: "2px 8px",
|
||||
borderRadius: 4,
|
||||
}}
|
||||
>
|
||||
{shortId(run.run_id)}
|
||||
</span>
|
||||
<span
|
||||
style={{
|
||||
fontSize: 11,
|
||||
color: "#a1a1aa",
|
||||
background: "#f4f4f5",
|
||||
padding: "2px 8px",
|
||||
borderRadius: 4,
|
||||
}}
|
||||
>
|
||||
{run.workflow_type}
|
||||
</span>
|
||||
</div>
|
||||
|
||||
{/* key fields grid */}
|
||||
<div
|
||||
style={{
|
||||
padding: "12px 20px",
|
||||
display: "grid",
|
||||
gridTemplateColumns: "repeat(auto-fill, minmax(220px, 1fr))",
|
||||
gap: "8px 24px",
|
||||
fontFamily: "monospace",
|
||||
fontSize: 12,
|
||||
}}
|
||||
>
|
||||
<FieldPair label="status">
|
||||
<span style={{ textTransform: "capitalize", color: "#27272a" }}>{run.status}</span>
|
||||
</FieldPair>
|
||||
<FieldPair label="created">
|
||||
<span style={{ color: "#27272a" }}>{timeAgo(run.created_at)}</span>
|
||||
</FieldPair>
|
||||
|
||||
{meta.pr_url && (
|
||||
<FieldPair label="pr">
|
||||
<a
|
||||
href={String(meta.pr_url)}
|
||||
target="_blank"
|
||||
rel="noopener noreferrer"
|
||||
style={{ color: "#2563eb", textDecoration: "none", wordBreak: "break-all" }}
|
||||
>
|
||||
{String(meta.pr_url)}
|
||||
</a>
|
||||
</FieldPair>
|
||||
)}
|
||||
|
||||
{primaryFields.map(({ key, label }) => {
|
||||
const v = meta[key];
|
||||
if (v === null || v === undefined || v === "") return null;
|
||||
const str = typeof v === "object" ? JSON.stringify(v) : String(v);
|
||||
return (
|
||||
<FieldPair key={key} label={label}>
|
||||
<TruncatedValue value={str} />
|
||||
</FieldPair>
|
||||
);
|
||||
})}
|
||||
|
||||
{extraEntries.map(([k, v]) => {
|
||||
const str = typeof v === "object" ? JSON.stringify(v) : String(v);
|
||||
return (
|
||||
<FieldPair key={k} label={k}>
|
||||
<TruncatedValue value={str} />
|
||||
</FieldPair>
|
||||
);
|
||||
})}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
const FieldPair: React.FC<{ label: string; children: React.ReactNode }> = ({
|
||||
label,
|
||||
children,
|
||||
}) => (
|
||||
<div style={{ display: "flex", flexDirection: "column", gap: 1 }}>
|
||||
<span style={{ fontSize: 10, color: "#a1a1aa", textTransform: "uppercase", letterSpacing: "0.06em" }}>
|
||||
{label}
|
||||
</span>
|
||||
<span style={{ fontSize: 12 }}>{children}</span>
|
||||
</div>
|
||||
);
|
||||
|
||||
// ── gantt timeline ────────────────────────────────────────────────────────────
|
||||
|
||||
const GanttTimeline: React.FC<{
|
||||
run: WorkflowRun;
|
||||
events: WorkflowRunEvent[];
|
||||
}> = ({ run, events }) => {
|
||||
if (events.length === 0) {
|
||||
return (
|
||||
<div style={{ padding: "16px 0", color: "#a1a1aa", fontSize: 12, fontFamily: "monospace" }}>
|
||||
No events recorded
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
const runStart = new Date(run.created_at).getTime();
|
||||
const eventTimes = events.map((e) => new Date(e.created_at).getTime());
|
||||
const lastTime = Math.max(...eventTimes);
|
||||
const totalSpan = Math.max(lastTime - runStart, 1);
|
||||
const totalDur = fmtDuration(lastTime - runStart);
|
||||
|
||||
return (
|
||||
<div style={{ fontFamily: "monospace", fontSize: 12 }}>
|
||||
{/* ruler */}
|
||||
<div style={{ display: "grid", gridTemplateColumns: "160px 1fr", gap: "0 12px", marginBottom: 2 }}>
|
||||
<div />
|
||||
<div style={{ position: "relative", height: 16 }}>
|
||||
{[0, 100].map((pct) => (
|
||||
<span
|
||||
key={pct}
|
||||
style={{
|
||||
position: "absolute",
|
||||
left: `${pct}%`,
|
||||
transform: pct === 100 ? "translateX(-100%)" : undefined,
|
||||
fontSize: 10,
|
||||
color: "#a1a1aa",
|
||||
}}
|
||||
>
|
||||
{pct === 0 ? "0" : totalDur}
|
||||
</span>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* outer run bar */}
|
||||
<div style={{ display: "grid", gridTemplateColumns: "160px 1fr", gap: "0 12px", marginBottom: 4 }}>
|
||||
<div style={{ color: "#3f3f46", overflow: "hidden", textOverflow: "ellipsis", whiteSpace: "nowrap", paddingTop: 2 }}>
|
||||
{runTitle(run)}
|
||||
</div>
|
||||
<div
|
||||
style={{
|
||||
height: 24,
|
||||
background: "#f4f4f5",
|
||||
border: "1px solid #d4d4d8",
|
||||
borderRadius: 4,
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
paddingLeft: 8,
|
||||
}}
|
||||
>
|
||||
<span style={{ color: "#71717a", fontSize: 11 }}>{totalDur}</span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* event rows */}
|
||||
<div style={{ display: "grid", gridTemplateColumns: "160px 1fr", gap: "0 12px", rowGap: 3 }}>
|
||||
{events.map((ev) => {
|
||||
const evTime = new Date(ev.created_at).getTime();
|
||||
const leftPct = ((evTime - runStart) / totalSpan) * 100;
|
||||
|
||||
const nextIdx = events.findIndex((e) => e.sequence_number > ev.sequence_number);
|
||||
const nextTime =
|
||||
nextIdx >= 0
|
||||
? new Date(events[nextIdx].created_at).getTime()
|
||||
: lastTime + Math.max(totalSpan * 0.12, 500);
|
||||
const widthPct = Math.max(8, ((nextTime - evTime) / totalSpan) * 100);
|
||||
const style = eventStyle(ev.event_type);
|
||||
const dur = fmtDuration(nextTime - evTime);
|
||||
|
||||
return (
|
||||
<React.Fragment key={ev.event_id}>
|
||||
<div
|
||||
style={{
|
||||
color: style.text,
|
||||
overflow: "hidden",
|
||||
textOverflow: "ellipsis",
|
||||
whiteSpace: "nowrap",
|
||||
paddingTop: 2,
|
||||
paddingLeft: 12,
|
||||
}}
|
||||
>
|
||||
{ev.step_name || ev.event_type}
|
||||
</div>
|
||||
<div style={{ position: "relative", height: 24 }}>
|
||||
<Tooltip
|
||||
title={
|
||||
<div style={{ fontFamily: "monospace", fontSize: 11, lineHeight: 1.6 }}>
|
||||
<div><span style={{ color: "#a1a1aa" }}>type: </span><span style={{ color: style.text }}>{ev.event_type}</span></div>
|
||||
<div><span style={{ color: "#a1a1aa" }}>step: </span>{ev.step_name}</div>
|
||||
<div><span style={{ color: "#a1a1aa" }}>seq: </span>{ev.sequence_number}</div>
|
||||
<div><span style={{ color: "#a1a1aa" }}>time: </span>{timeAgo(ev.created_at)}</div>
|
||||
{ev.data && Object.keys(ev.data).length > 0 && (
|
||||
<div><span style={{ color: "#a1a1aa" }}>data: </span>{JSON.stringify(ev.data)}</div>
|
||||
)}
|
||||
</div>
|
||||
}
|
||||
>
|
||||
<div
|
||||
style={{
|
||||
position: "absolute",
|
||||
left: `${Math.min(leftPct, 92)}%`,
|
||||
width: `${Math.min(widthPct, 100 - Math.min(leftPct, 92))}%`,
|
||||
height: "100%",
|
||||
background: style.bar,
|
||||
border: `1px solid ${style.border}`,
|
||||
borderRadius: 4,
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
paddingLeft: 8,
|
||||
cursor: "default",
|
||||
overflow: "hidden",
|
||||
gap: 6,
|
||||
}}
|
||||
>
|
||||
<span style={{ color: style.text, whiteSpace: "nowrap", fontSize: 11 }}>{ev.event_type}</span>
|
||||
{dur && <span style={{ color: "#a1a1aa", whiteSpace: "nowrap", fontSize: 11 }}>{dur}</span>}
|
||||
</div>
|
||||
</Tooltip>
|
||||
</div>
|
||||
</React.Fragment>
|
||||
);
|
||||
})}
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
// ── message row ───────────────────────────────────────────────────────────────
|
||||
|
||||
const MessageRow: React.FC<{ msg: WorkflowRunMessage }> = ({ msg }) => {
|
||||
const roleColor: Record<string, string> = {
|
||||
user: "#2563eb",
|
||||
assistant: "#16a34a",
|
||||
system: "#7c3aed",
|
||||
tool_result: "#d97706",
|
||||
};
|
||||
const color = roleColor[msg.role] ?? "#52525b";
|
||||
|
||||
return (
|
||||
<div
|
||||
style={{
|
||||
display: "grid",
|
||||
gridTemplateColumns: "80px 1fr",
|
||||
gap: "0 16px",
|
||||
padding: "10px 0",
|
||||
borderBottom: "1px solid #f4f4f5",
|
||||
fontFamily: "monospace",
|
||||
fontSize: 12,
|
||||
alignItems: "start",
|
||||
}}
|
||||
>
|
||||
<span style={{ color, paddingTop: 1 }}>[{msg.role}]</span>
|
||||
<div>
|
||||
<span
|
||||
style={{
|
||||
color: "#27272a",
|
||||
lineHeight: 1.6,
|
||||
whiteSpace: "pre-wrap",
|
||||
wordBreak: "break-word",
|
||||
display: "block",
|
||||
}}
|
||||
>
|
||||
{msg.content}
|
||||
</span>
|
||||
<span style={{ color: "#a1a1aa", fontSize: 11, marginTop: 2, display: "block" }}>
|
||||
{timeAgo(msg.created_at)}
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
// ── main component ────────────────────────────────────────────────────────────
|
||||
|
||||
const WorkflowRuns: React.FC<WorkflowRunsProps> = ({ accessToken }) => {
|
||||
const [runs, setRuns] = useState<WorkflowRun[]>([]);
|
||||
const [loadingRuns, setLoadingRuns] = useState(false);
|
||||
const [selectedRun, setSelectedRun] = useState<WorkflowRun | null>(null);
|
||||
const [events, setEvents] = useState<WorkflowRunEvent[]>([]);
|
||||
const [messages, setMessages] = useState<WorkflowRunMessage[]>([]);
|
||||
const [loadingDetail, setLoadingDetail] = useState(false);
|
||||
const [drawerOpen, setDrawerOpen] = useState(false);
|
||||
|
||||
const fetchRuns = useCallback(async () => {
|
||||
if (!accessToken) return;
|
||||
setLoadingRuns(true);
|
||||
try {
|
||||
const res = await fetch(`${proxyBaseUrl ?? ""}/v1/workflows/runs?limit=100`, {
|
||||
headers: { Authorization: `Bearer ${accessToken}` },
|
||||
});
|
||||
if (!res.ok) throw new Error(`HTTP ${res.status}`);
|
||||
const data = await res.json();
|
||||
setRuns(data.runs ?? []);
|
||||
} catch (err) {
|
||||
console.error("workflow runs fetch failed:", err);
|
||||
} finally {
|
||||
setLoadingRuns(false);
|
||||
}
|
||||
}, [accessToken]);
|
||||
|
||||
const fetchRunDetail = useCallback(
|
||||
async (run: WorkflowRun) => {
|
||||
if (!accessToken) return;
|
||||
setSelectedRun(run);
|
||||
setDrawerOpen(true);
|
||||
setLoadingDetail(true);
|
||||
setEvents([]);
|
||||
setMessages([]);
|
||||
try {
|
||||
const base = proxyBaseUrl ?? "";
|
||||
const [evRes, msgRes] = await Promise.all([
|
||||
fetch(`${base}/v1/workflows/runs/${run.run_id}/events`, {
|
||||
headers: { Authorization: `Bearer ${accessToken}` },
|
||||
}),
|
||||
fetch(`${base}/v1/workflows/runs/${run.run_id}/messages`, {
|
||||
headers: { Authorization: `Bearer ${accessToken}` },
|
||||
}),
|
||||
]);
|
||||
const evData = evRes.ok ? await evRes.json() : { events: [] };
|
||||
const msgData = msgRes.ok ? await msgRes.json() : { messages: [] };
|
||||
setEvents(
|
||||
[...(evData.events ?? [])].sort(
|
||||
(a: WorkflowRunEvent, b: WorkflowRunEvent) => a.sequence_number - b.sequence_number
|
||||
)
|
||||
);
|
||||
setMessages(
|
||||
[...(msgData.messages ?? [])].sort(
|
||||
(a: WorkflowRunMessage, b: WorkflowRunMessage) => a.sequence_number - b.sequence_number
|
||||
)
|
||||
);
|
||||
} catch (err) {
|
||||
console.error("workflow run detail fetch failed:", err);
|
||||
} finally {
|
||||
setLoadingDetail(false);
|
||||
}
|
||||
},
|
||||
[accessToken]
|
||||
);
|
||||
|
||||
useEffect(() => {
|
||||
fetchRuns();
|
||||
}, [fetchRuns]);
|
||||
|
||||
const columns = [
|
||||
{
|
||||
title: "Run",
|
||||
dataIndex: "run_id",
|
||||
key: "run",
|
||||
render: (_: string, run: WorkflowRun) => (
|
||||
<div style={{ display: "flex", alignItems: "center", gap: 8 }}>
|
||||
<StatusDot status={run.status} size={7} />
|
||||
<div>
|
||||
<div style={{ fontSize: 13, color: "#18181b", fontWeight: 500, lineHeight: 1.4 }}>
|
||||
{runTitle(run)}
|
||||
</div>
|
||||
<div style={{ fontFamily: "monospace", fontSize: 11, color: "#a1a1aa" }}>
|
||||
{shortId(run.run_id)}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
),
|
||||
},
|
||||
{
|
||||
title: "Type",
|
||||
dataIndex: "workflow_type",
|
||||
key: "workflow_type",
|
||||
render: (v: string) => (
|
||||
<span style={{ fontFamily: "monospace", fontSize: 12, color: "#71717a" }}>{v}</span>
|
||||
),
|
||||
},
|
||||
{
|
||||
title: "Status",
|
||||
dataIndex: "status",
|
||||
key: "status",
|
||||
render: (status: RunStatus, run: WorkflowRun) => {
|
||||
const state = run.metadata?.state;
|
||||
return (
|
||||
<div style={{ display: "flex", alignItems: "center", gap: 6 }}>
|
||||
<StatusDot status={status} size={7} />
|
||||
<span style={{ fontSize: 12, color: "#52525b", textTransform: "capitalize" }}>
|
||||
{state ?? status}
|
||||
</span>
|
||||
</div>
|
||||
);
|
||||
},
|
||||
},
|
||||
{
|
||||
title: "Created",
|
||||
dataIndex: "created_at",
|
||||
key: "created_at",
|
||||
render: (v: string) => (
|
||||
<span style={{ fontSize: 12, color: "#a1a1aa" }}>{timeAgo(v)}</span>
|
||||
),
|
||||
},
|
||||
];
|
||||
|
||||
return (
|
||||
<div
|
||||
style={{
|
||||
padding: "24px 32px",
|
||||
fontFamily: '-apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, sans-serif',
|
||||
minHeight: "calc(100vh - 64px)",
|
||||
background: "#fff",
|
||||
}}
|
||||
>
|
||||
{/* page header */}
|
||||
<div
|
||||
style={{
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
justifyContent: "space-between",
|
||||
marginBottom: 20,
|
||||
}}
|
||||
>
|
||||
<div>
|
||||
<div style={{ fontSize: 18, fontWeight: 600, color: "#18181b" }}>Workflow Runs</div>
|
||||
<div style={{ fontSize: 13, color: "#71717a", marginTop: 2 }}>
|
||||
Durable state tracking for agents and automated workflows
|
||||
</div>
|
||||
</div>
|
||||
<Button
|
||||
icon={<ReloadOutlined />}
|
||||
onClick={fetchRuns}
|
||||
loading={loadingRuns}
|
||||
style={{ color: "#71717a", borderColor: "#e4e4e7" }}
|
||||
>
|
||||
Refresh
|
||||
</Button>
|
||||
</div>
|
||||
|
||||
{/* runs table — matches logs page density */}
|
||||
<div className="rounded-lg custom-border overflow-x-auto w-full">
|
||||
<Table
|
||||
dataSource={runs}
|
||||
columns={columns}
|
||||
rowKey="run_id"
|
||||
loading={loadingRuns}
|
||||
size="small"
|
||||
pagination={{ pageSize: 50, hideOnSinglePage: true, size: "small" }}
|
||||
onRow={(run) => ({
|
||||
onClick: () => fetchRunDetail(run),
|
||||
style: { cursor: "pointer" },
|
||||
})}
|
||||
locale={{
|
||||
emptyText: (
|
||||
<Empty
|
||||
description={<span style={{ color: "#a1a1aa", fontSize: 13 }}>No workflow runs yet</span>}
|
||||
image={Empty.PRESENTED_IMAGE_SIMPLE}
|
||||
/>
|
||||
),
|
||||
}}
|
||||
className="[&_.ant-table-cell]:py-0.5 [&_.ant-table-thead_.ant-table-cell]:py-1"
|
||||
style={{ border: "none" }}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{/* detail drawer */}
|
||||
<Drawer
|
||||
open={drawerOpen}
|
||||
onClose={() => setDrawerOpen(false)}
|
||||
width={680}
|
||||
title={null}
|
||||
closable={false}
|
||||
bodyStyle={{ padding: 0 }}
|
||||
styles={{ body: { padding: 0 } }}
|
||||
>
|
||||
{!selectedRun ? null : loadingDetail ? (
|
||||
<div style={{ display: "flex", justifyContent: "center", padding: 80 }}>
|
||||
<Spin />
|
||||
</div>
|
||||
) : (
|
||||
<div style={{ padding: "24px 28px", fontFamily: '-apple-system, BlinkMacSystemFont, "Segoe UI", Roboto, sans-serif' }}>
|
||||
{/* drawer close + refresh */}
|
||||
<div
|
||||
style={{
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
justifyContent: "space-between",
|
||||
marginBottom: 16,
|
||||
}}
|
||||
>
|
||||
<button
|
||||
onClick={() => setDrawerOpen(false)}
|
||||
style={{
|
||||
background: "none",
|
||||
border: "none",
|
||||
cursor: "pointer",
|
||||
padding: "4px 0",
|
||||
fontSize: 12,
|
||||
color: "#a1a1aa",
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
gap: 4,
|
||||
}}
|
||||
>
|
||||
← close
|
||||
</button>
|
||||
<Button
|
||||
size="small"
|
||||
icon={<ReloadOutlined />}
|
||||
onClick={() => fetchRunDetail(selectedRun)}
|
||||
loading={loadingDetail}
|
||||
style={{ color: "#71717a", borderColor: "#e4e4e7" }}
|
||||
>
|
||||
Refresh
|
||||
</Button>
|
||||
</div>
|
||||
|
||||
{/* metadata card — top */}
|
||||
<MetadataCard run={selectedRun} />
|
||||
|
||||
{/* collapsible sections */}
|
||||
<Collapse
|
||||
defaultActiveKey={["timeline"]}
|
||||
ghost={false}
|
||||
style={{ border: "1px solid #e4e4e7", borderRadius: 8, overflow: "hidden" }}
|
||||
items={[
|
||||
{
|
||||
key: "timeline",
|
||||
label: (
|
||||
<span style={{ fontSize: 12, fontWeight: 500, color: "#3f3f46" }}>
|
||||
Timeline
|
||||
<span style={{ marginLeft: 6, fontSize: 11, color: "#a1a1aa", fontWeight: 400 }}>
|
||||
{events.length} {events.length === 1 ? "event" : "events"}
|
||||
</span>
|
||||
</span>
|
||||
),
|
||||
children: (
|
||||
<div style={{ padding: "4px 4px 12px" }}>
|
||||
<GanttTimeline run={selectedRun} events={events} />
|
||||
</div>
|
||||
),
|
||||
},
|
||||
{
|
||||
key: "messages",
|
||||
label: (
|
||||
<span style={{ fontSize: 12, fontWeight: 500, color: "#3f3f46" }}>
|
||||
Messages
|
||||
<span style={{ marginLeft: 6, fontSize: 11, color: "#a1a1aa", fontWeight: 400 }}>
|
||||
{messages.length}
|
||||
</span>
|
||||
</span>
|
||||
),
|
||||
children: messages.length === 0 ? (
|
||||
<div style={{ padding: "12px 4px", color: "#a1a1aa", fontSize: 12, fontFamily: "monospace" }}>
|
||||
No messages
|
||||
</div>
|
||||
) : (
|
||||
<div style={{ paddingBottom: 4 }}>
|
||||
{messages.map((msg) => (
|
||||
<MessageRow key={msg.message_id} msg={msg} />
|
||||
))}
|
||||
</div>
|
||||
),
|
||||
},
|
||||
]}
|
||||
/>
|
||||
</div>
|
||||
)}
|
||||
</Drawer>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
export default WorkflowRuns;
|
||||
Loading…
Add table
Reference in a new issue