Merge branch 'litellm_internal_staging' into litellm_lit_4116_drop_params_string_coerce

This commit is contained in:
mateo-berri 2026-09-07 16:22:26 -07:00
commit b827375e60
18 changed files with 1058 additions and 69 deletions

View file

@ -2002,6 +2002,9 @@ if TYPE_CHECKING:
from .llms.hosted_vllm.responses.transformation import (
HostedVLLMResponsesAPIConfig as HostedVLLMResponsesAPIConfig,
)
from .llms.fireworks_ai.responses.transformation import (
FireworksAIResponsesAPIConfig as FireworksAIResponsesAPIConfig,
)
from .llms.github_copilot.chat.transformation import (
GithubCopilotConfig as GithubCopilotConfig,
)

View file

@ -237,6 +237,7 @@ LLM_CONFIG_NAMES: Final = (
"XAIResponsesAPIConfig",
"LiteLLMProxyResponsesAPIConfig",
"HostedVLLMResponsesAPIConfig",
"FireworksAIResponsesAPIConfig",
"VolcEngineResponsesAPIConfig",
"PerplexityResponsesConfig",
"DatabricksResponsesAPIConfig",
@ -957,6 +958,10 @@ _LLM_CONFIGS_IMPORT_MAP: Final = {
".llms.hosted_vllm.responses.transformation",
"HostedVLLMResponsesAPIConfig",
),
"FireworksAIResponsesAPIConfig": (
".llms.fireworks_ai.responses.transformation",
"FireworksAIResponsesAPIConfig",
),
"VolcEngineResponsesAPIConfig": (
".llms.volcengine.responses.transformation",
"VolcEngineResponsesAPIConfig",

View file

@ -1,3 +1,5 @@
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final
from httpx import Headers
@ -13,7 +15,7 @@ class FireworksAIException(BaseLLMException):
pass
def get_fireworks_session_id(litellm_params: dict) -> str | None:
def get_fireworks_session_id(litellm_params: Mapping[str, object]) -> str | None:
"""
Session id to send as `x-session-affinity`, or None when the caller gave none.
@ -23,19 +25,39 @@ def get_fireworks_session_id(litellm_params: dict) -> str | None:
"""
params: Final = litellm_params
metadata: Final = params.get("metadata")
if isinstance(metadata, dict) and metadata.get(SESSION_ID_GENERATED_METADATA_KEY):
if isinstance(metadata, Mapping) and metadata.get(SESSION_ID_GENERATED_METADATA_KEY):
return None
for key in ("litellm_session_id", "session_id"):
value = params.get(key)
if value:
return str(value)
if isinstance(metadata, dict):
if isinstance(metadata, Mapping):
value = metadata.get("session_id")
if value:
return str(value)
return None
def with_fireworks_session_affinity(
headers: Mapping[str, str], litellm_params: Mapping[str, object]
) -> Mapping[str, str]:
if any(key.lower() == "x-session-affinity" for key in headers):
return headers
session_id: Final = get_fireworks_session_id(litellm_params)
if not session_id:
return headers
return MappingProxyType({**headers, "x-session-affinity": session_id})
def resolve_fireworks_api_key(api_key: str | None) -> str | None:
return api_key or (
get_secret_str("FIREWORKS_API_KEY")
or get_secret_str("FIREWORKS_AI_API_KEY")
or get_secret_str("FIREWORKSAI_API_KEY")
or get_secret_str("FIREWORKS_AI_TOKEN")
)
AZURE_FOUNDRY_FIREWORKS_MODEL_ID_PREFIX: Final = "FW-"
@ -63,13 +85,7 @@ class FireworksAIMixin:
)
def _get_api_key(self, api_key: str | None) -> str | None:
dynamic_api_key: Final = api_key or (
get_secret_str("FIREWORKS_API_KEY")
or get_secret_str("FIREWORKS_AI_API_KEY")
or get_secret_str("FIREWORKSAI_API_KEY")
or get_secret_str("FIREWORKS_AI_TOKEN")
)
return dynamic_api_key
return resolve_fireworks_api_key(api_key)
def validate_environment(
self,
@ -92,9 +108,5 @@ class FireworksAIMixin:
return self._add_session_affinity_header({**auth_headers, **content_type_header}, litellm_params)
def _add_session_affinity_header(self, headers: dict, litellm_params: dict) -> dict:
if any(key.lower() == "x-session-affinity" for key in headers):
return headers
session_id: Final = get_fireworks_session_id(litellm_params)
if not session_id:
return headers
return {**headers, "x-session-affinity": session_id}
pinned: Final = with_fireworks_session_affinity(headers, litellm_params)
return dict(pinned) # mutable-ok: the HTTP handler updates the returned headers in place

View file

@ -0,0 +1,102 @@
from collections.abc import Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
from urllib.parse import unquote
import httpx
from openai.types.responses import EasyInputMessageParam, ResponseInputItemParam
from litellm.llms.fireworks_ai.common_utils import (
resolve_fireworks_api_key,
resolve_fireworks_resource_name,
with_fireworks_session_affinity,
)
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import ResponseInputParam
from litellm.types.responses.main import DeleteResponseResult
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
FIREWORKS_AI_DEFAULT_API_BASE: Final = "https://api.fireworks.ai/inference/v1"
def _session_params(litellm_params: GenericLiteLLMParams) -> Mapping[str, object]:
extras: Final[Mapping[str, object]] = litellm_params.model_extra or MappingProxyType({})
return MappingProxyType(
{"litellm_session_id": extras.get("litellm_session_id"), "metadata": extras.get("litellm_metadata")}
)
def _developer_item_as_system(item: ResponseInputItemParam) -> ResponseInputItemParam:
if "role" not in item or item["role"] != "developer":
return item
return EasyInputMessageParam(role="system", content=item["content"], type="message")
def _developer_items_as_system(input: str | ResponseInputParam) -> str | ResponseInputParam:
if isinstance(input, str):
return input
return [_developer_item_as_system(item) for item in input]
class FireworksAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
@property
def custom_llm_provider(self) -> LlmProviders:
return LlmProviders.FIREWORKS_AI
def validate_environment(
self,
headers: Mapping[str, str],
model: str,
litellm_params: GenericLiteLLMParams | None,
) -> dict: # mutable-ok: overrides the base class signature
params: Final = litellm_params or GenericLiteLLMParams()
api_key: Final = resolve_fireworks_api_key(params.api_key)
if api_key is None:
raise ValueError("FIREWORKS_API_KEY is not set")
authorized: Final = MappingProxyType(
{"Content-Type": "application/json", **headers, "Authorization": f"Bearer {api_key}"}
)
pinned: Final = with_fireworks_session_affinity(authorized, _session_params(params))
return dict(pinned) # mutable-ok: the HTTP handler updates the returned headers in place
def get_complete_url(self, api_base: str | None, litellm_params: Mapping[str, object]) -> str:
base: Final = (api_base or get_secret_str("FIREWORKS_API_BASE") or FIREWORKS_AI_DEFAULT_API_BASE).rstrip("/")
return f"{base}/responses"
def _validate_input_param(self, input: str | ResponseInputParam) -> str | ResponseInputParam:
return _developer_items_as_system(super()._validate_input_param(input))
def transform_responses_api_request(
self,
model: str,
input: str | ResponseInputParam,
response_api_optional_request_params: dict, # mutable-ok: overrides the base class signature
litellm_params: GenericLiteLLMParams,
headers: dict, # mutable-ok: overrides the base class signature
) -> dict: # mutable-ok: overrides the base class signature
return super().transform_responses_api_request(
model=resolve_fireworks_resource_name(model),
input=input,
response_api_optional_request_params=response_api_optional_request_params,
litellm_params=litellm_params,
headers=headers,
)
def transform_delete_response_api_response(
self,
raw_response: httpx.Response,
logging_obj: "LiteLLMLoggingObj",
) -> DeleteResponseResult:
deleted_id: Final = unquote(raw_response.request.url.path.rsplit("/", 1)[-1])
return DeleteResponseResult(id=deleted_id, object="response", deleted=True)
def supports_native_websocket(self) -> bool:
return False
def supports_native_file_search(self) -> bool:
return False

View file

@ -6,10 +6,12 @@ Canonical definition for ``litellm_mcpservertable``. Re-exported from
"""
import enum
from collections.abc import Mapping
from datetime import datetime
from types import MappingProxyType
from typing import Literal
from pydantic import Field
from pydantic import Field, ValidationInfo, field_validator
from litellm.types.llms.base import LiteLLMPydanticObjectBase
from litellm.types.mcp import MCPAuthType, MCPCredentials, MCPTransportType
@ -115,3 +117,12 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
submitted_at: datetime | None = None
reviewed_at: datetime | None = None
review_notes: str | None = None
@field_validator("static_headers", "env", mode="before")
@classmethod
def decode_stored_secret_map(cls, value: object, info: ValidationInfo) -> Mapping[str, str] | None:
from litellm.proxy.common_utils.encrypt_decrypt_utils import decode_secret_map
if value is None and info.field_name == "env":
return MappingProxyType({})
return decode_secret_map(value, key=info.field_name or "secret map")

View file

@ -23,8 +23,11 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
SecretMapDecodeError,
_get_salt_key,
decode_secret_map,
decrypt_value_helper,
encrypt_secret_map,
encrypt_value_helper,
)
from litellm.proxy.utils import PrismaClient
@ -360,7 +363,7 @@ def _prepare_mcp_server_data(
# exclude_unset filter is respected. Reading back from ``data`` would
# reintroduce defaults (e.g. ``env={}``) for fields the caller never set.
if data_dict.get("static_headers") is not None:
data_dict["static_headers"] = safe_dumps(data_dict["static_headers"])
data_dict["static_headers"] = encrypt_secret_map(data_dict["static_headers"])
# env_vars is read from ``data_dict`` (not ``data``) like every other JSON
# column so the exclude_unset filter is respected: a partial update that
@ -376,7 +379,7 @@ def _prepare_mcp_server_data(
data_dict["mcp_info"] = safe_dumps(data_dict["mcp_info"])
if data_dict.get("env") is not None:
data_dict["env"] = safe_dumps(data_dict["env"])
data_dict["env"] = encrypt_secret_map(data_dict["env"])
if "tool_name_to_display_name" in data_dict:
data_dict["tool_name_to_display_name"] = safe_dumps(data_dict["tool_name_to_display_name"] or {})
@ -589,6 +592,19 @@ def decrypt_credentials(
return credentials
def _readable_mcp_servers(
rows: Iterable["prisma_db_models.LiteLLM_MCPServerTable"],
) -> Iterable[LiteLLM_MCPServerTable]:
for row in rows:
try:
table = LiteLLM_MCPServerTable.model_validate(row.model_dump())
except SecretMapDecodeError:
verbose_proxy_logger.warning("Skipping MCP server %s: cannot decrypt secret map", row.server_id)
continue
decrypt_global_env_var_values(table.env_vars)
yield table
async def get_all_mcp_servers(
prisma_client: PrismaClient,
approval_status: str | None = None,
@ -609,10 +625,7 @@ async def get_all_mcp_servers(
)
mcp_servers: Final = await _db_find_mcp_server_rows(prisma_client, where)
tables: Final = [LiteLLM_MCPServerTable.model_validate(mcp_server.model_dump()) for mcp_server in mcp_servers]
for table in tables:
decrypt_global_env_var_values(table.env_vars)
return tables
return list(_readable_mcp_servers(mcp_servers))
async def get_mcp_server(prisma_client: PrismaClient, server_id: str) -> LiteLLM_MCPServerTable | None:
@ -638,13 +651,7 @@ async def get_mcp_servers(prisma_client: PrismaClient, server_ids: Iterable[str]
"server_id": {"in": server_ids},
}
)
final_mcp_servers: Final[list[LiteLLM_MCPServerTable]] = []
for _mcp_server in _mcp_servers:
table = LiteLLM_MCPServerTable.model_validate(_mcp_server.model_dump())
decrypt_global_env_var_values(table.env_vars)
final_mcp_servers.append(table)
return final_mcp_servers
return list(_readable_mcp_servers(_mcp_servers))
async def get_mcp_servers_by_verificationtoken(prisma_client: PrismaClient, token: str) -> list[str]:
@ -852,12 +859,10 @@ async def create_mcp_server(
data_dict["created_by"] = touched_by
data_dict["updated_by"] = touched_by
new_mcp_server: Final[LiteLLM_MCPServerTable] = await MCPServerRepository(prisma_client).table.create(
data=data_dict, # pyright: ignore[reportAssignmentType] # prisma row, not domain LiteLLM_MCPServerTable
)
new_mcp_server: Final = await MCPServerRepository(prisma_client).table.create(data=data_dict)
_decrypt_env_vars_on_returned_row(new_mcp_server)
return new_mcp_server
return LiteLLM_MCPServerTable.model_validate(new_mcp_server.model_dump())
async def create_draft_mcp_server(
@ -1066,13 +1071,13 @@ async def update_mcp_server(
data_dict["credentials"] = Json(None)
updated_mcp_server: Final[LiteLLM_MCPServerTable | None] = await MCPServerRepository(prisma_client).table.update(
updated_mcp_server: Final = await MCPServerRepository(prisma_client).table.update(
where={"server_id": data.server_id},
data=data_dict, # pyright: ignore[reportAssignmentType] # prisma row, not domain LiteLLM_MCPServerTable
data=data_dict,
)
_decrypt_env_vars_on_returned_row(updated_mcp_server)
return updated_mcp_server
return LiteLLM_MCPServerTable.model_validate(updated_mcp_server.model_dump()) if updated_mcp_server else None
async def get_mcp_server_oauth_client_credentials(prisma_client: PrismaClient, server_id: str) -> object | None:
@ -1144,6 +1149,13 @@ async def rotate_mcp_server_credentials_master_key(prisma_client: PrismaClient,
if rotated_env_vars is not None:
update_data["env_vars"] = safe_dumps(rotated_env_vars)
for field in ("static_headers", "env"):
try:
if secret_map := decode_secret_map(getattr(mcp_server, field, None), key=field):
update_data[field] = encrypt_secret_map(secret_map, new_encryption_key=new_master_key)
except SecretMapDecodeError:
verbose_proxy_logger.warning("Cannot rotate MCP %s for server %s", field, mcp_server.server_id)
if not update_data:
continue
@ -1894,9 +1906,7 @@ async def get_mcp_submissions(
order={"submitted_at": "desc"},
take=500, # safety cap; paginate if needed in a future iteration
)
items: Final = [LiteLLM_MCPServerTable.model_validate(r.model_dump()) for r in rows]
for item in items:
decrypt_global_env_var_values(item.env_vars)
items: Final = list(_readable_mcp_servers(rows))
pending: Final = sum(1 for i in items if i.approval_status == MCPApprovalStatus.pending_review)
active: Final = sum(1 for i in items if i.approval_status == MCPApprovalStatus.active)

View file

@ -6272,8 +6272,7 @@ class MCPServerManager:
]
}
)
db_mcp_servers: Final = [LiteLLM_MCPServerTable.model_validate(r.model_dump()) for r in raw_rows]
verbose_logger.info("Found %s MCP servers in database", len(db_mcp_servers))
verbose_logger.info("Found %s MCP servers in database", len(raw_rows))
previous_registry: Final = self.registry
new_registry: Final[dict[str, MCPServer]] = {}
@ -6281,8 +6280,9 @@ class MCPServerManager:
# 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:
for row in raw_rows:
try:
server = LiteLLM_MCPServerTable.model_validate(row.model_dump())
existing_server = previous_registry.get(server.server_id)
if (
@ -6320,8 +6320,8 @@ class MCPServerManager:
except Exception as e:
verbose_logger.exception(
"Skipping MCP server %s (%s) during DB reload: %s",
server.server_id,
getattr(server, "alias", None),
getattr(row, "server_id", None),
getattr(row, "alias", None),
e,
)

View file

@ -1,7 +1,10 @@
import base64
import os
from collections.abc import Mapping
from typing import Final, Literal, cast
from pydantic import TypeAdapter, ValidationError
from litellm._logging import verbose_proxy_logger
# Versioned ciphertext marker for AES-256-GCM values.
@ -203,3 +206,40 @@ def decrypt_value(value: bytes, signing_key: str) -> str:
return plaintext
except Exception as e:
raise e
class SecretMapDecodeError(RuntimeError):
pass
_SECRET_MAP: Final = TypeAdapter(Mapping[str, str])
_STORED_SECRET_MAP: Final = TypeAdapter(Mapping[str, str] | str)
_SECRET_STRING: Final = TypeAdapter(str)
def encrypt_secret_map(value: Mapping[str, str], new_encryption_key: str | None = None) -> str:
if not value:
return "{}"
ciphertext: Final = _SECRET_STRING.validate_python(
encrypt_value_helper(_SECRET_MAP.dump_json(value).decode(), new_encryption_key=new_encryption_key), strict=True
)
return _SECRET_STRING.dump_json(ciphertext).decode()
def decode_secret_map(value: object, *, key: str) -> Mapping[str, str] | None:
if value is None:
return None
try:
stored: Final = (
_STORED_SECRET_MAP.validate_json(value, strict=True)
if isinstance(value, str) and value.lstrip().startswith(("{", '"'))
else _STORED_SECRET_MAP.validate_python(value, strict=True)
)
if not isinstance(stored, str):
return stored
decrypted: Final = decrypt_value_helper(
value=stored, key=key, exception_type="debug", return_original_value=False
)
return _SECRET_MAP.validate_json(decrypted, strict=True)
except ValidationError:
raise SecretMapDecodeError(f"Cannot decode encrypted MCP {key}; check LITELLM_SALT_KEY") from None

View file

@ -43,7 +43,9 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
_ALGO_AES_GCM,
_ENCRYPTION_ALGORITHM_SETTING,
_V2_GCM_PREFIX,
SecretMapDecodeError,
_get_salt_key,
decode_secret_map,
decrypt_value_helper,
encrypt_value_helper,
)
@ -65,6 +67,20 @@ class LocationReport:
# Used by --check (read-only classification):
legacy: int = 0 # nacl ciphertext still awaiting migration
def count(self, classification: ValueClass | None) -> None:
if classification is None:
return
self.scanned += 1
match classification:
case "migrated":
self.already_v2 += 1
case "legacy":
self.legacy += 1
case "undecryptable":
self.undecryptable += 1
case _:
self.plaintext += 1
def as_dict(self) -> dict[str, int]:
return {
"scanned": self.scanned,
@ -441,7 +457,7 @@ def _classify_callback_value(value: object) -> ValueClass:
_COVERED_TABLE_SPECS: Final = [
("model_table", "litellm_proxymodeltable", ("litellm_params",), ()),
("credentials", "litellm_credentialstable", ("credential_values",), ()),
("mcp_server", "litellm_mcpservertable", ("credentials", "env_vars"), ()),
("mcp_server", "litellm_mcpservertable", ("credentials", "env_vars", "static_headers", "env"), ()),
("mcp_user_credentials", "litellm_mcpusercredentials", (), ("credential_b64",)),
("mcp_user_env_vars", "litellm_mcpuserenvvars", (), ("values_b64",)),
]
@ -472,14 +488,18 @@ def _classify_into_report(report: LocationReport, value: str) -> None:
names, base URLs, ) do not decrypt and fall through to ``plaintext``, so
over-scanning a column is harmless to the residual count.
"""
report.scanned += 1
cls: Final = classify_value(value, key="scan")
if cls == "migrated":
report.already_v2 += 1
elif cls == "legacy":
report.legacy += 1
else: # plaintext / not-a-string
report.plaintext += 1
report.count(classify_value(value, key="scan"))
def _classify_secret_map(value: object, key: str) -> ValueClass | None:
try:
decoded: Final = decode_secret_map(value, key=key)
except SecretMapDecodeError:
return "undecryptable"
if not decoded:
return None
ciphertext: Final = json.loads(value) if isinstance(value, str) and value.lstrip().startswith('"') else value
return "migrated" if is_migrated(ciphertext) else "legacy"
async def _scan_one_table(
@ -503,6 +523,9 @@ async def _scan_one_table(
raw = getattr(row, col, None)
if raw is None:
continue
if db_attr == "litellm_mcpservertable" and col in ("static_headers", "env"):
report.count(_classify_secret_map(raw, col))
continue
if isinstance(raw, str):
try:
raw = json.loads(raw)

View file

@ -8699,6 +8699,8 @@ class ProviderConfigManager:
return litellm.OpenRouterResponsesAPIConfig()
elif litellm.LlmProviders.HOSTED_VLLM == provider:
return litellm.HostedVLLMResponsesAPIConfig()
elif litellm.LlmProviders.FIREWORKS_AI == provider:
return litellm.FireworksAIResponsesAPIConfig()
elif litellm.LlmProviders.BEDROCK_MANTLE == provider:
# Both decisions are data-driven from the model's price-map entry, with
# no model-name logic. Capability (can it serve Responses?) comes from

View file

@ -0,0 +1,455 @@
import json
from collections.abc import Mapping
from types import MappingProxyType
from typing import Final, TypedDict, cast
from unittest.mock import MagicMock, patch
from urllib.parse import quote
import httpx
import pytest
from openai.types.responses import (
EasyInputMessage,
ResponseFunctionToolCall,
ResponseOutputMessage,
ResponseOutputText,
ResponseReasoningItem,
)
from openai.types.responses.response_input_param import FunctionCallOutput
from openai.types.responses.tool_param import Mcp
from typing_extensions import ReadOnly
import litellm
from litellm.llms.fireworks_ai.responses.transformation import FireworksAIResponsesAPIConfig
from litellm.responses.file_search.emulated_handler import should_use_emulated_file_search
from litellm.types.llms.openai import InputTokensDetails, ResponseAPIUsage, ResponseInputParam, ResponsesAPIResponse
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
FIREWORKS_RESPONSES_URL: Final = "https://api.fireworks.ai/inference/v1/responses"
HTTPX_CLIENT_FACTORY: Final = "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client"
NO_HEADERS: Final[Mapping[str, str]] = MappingProxyType({})
NO_PARAMS: Final[Mapping[str, object]] = MappingProxyType({})
class _GeneratedSessionMetadata(TypedDict):
litellm_session_id_generated: ReadOnly[bool]
def _fireworks_response(model: str) -> Mapping[str, object]:
return ResponsesAPIResponse(
id="resp_0e946f2d46bf4b49bf8b29ff78083583",
object="response",
created_at=1788550000,
model=model,
status="completed",
output=(
ResponseReasoningItem(id="rs_1", summary=(), type="reasoning"),
ResponseOutputMessage(
id="msg_1",
status="completed",
role="assistant",
type="message",
content=(ResponseOutputText(type="output_text", text="Paris is clear and 21C.", annotations=()),),
),
ResponseFunctionToolCall(
id="fc_1",
call_id="call_abc123",
name="get_weather",
arguments='{"city": "Paris"}',
status="completed",
type="function_call",
),
),
usage=ResponseAPIUsage(
input_tokens=179,
output_tokens=100,
total_tokens=279,
input_tokens_details=InputTokensDetails(cached_tokens=0),
),
).model_dump(mode="json", exclude_none=True)
def _mock_http_client(response_body: Mapping[str, object]) -> MagicMock:
client: Final = MagicMock()
response: Final = MagicMock()
response.status_code = 200
response.headers = httpx.Headers((("content-type", "application/json"),))
response.json.return_value = response_body
response.text = json.dumps(response_body)
client.post.return_value = response
return client
def _sent_request(client: MagicMock) -> tuple[str, Mapping[str, str], Mapping[str, object]]:
kwargs: Final = client.post.call_args.kwargs
body: Final = kwargs["json"] if "json" in kwargs else json.loads(kwargs["data"])
return kwargs["url"], kwargs["headers"], body
@pytest.fixture(autouse=True)
def fireworks_env(monkeypatch: pytest.MonkeyPatch) -> None:
for name in (
"FIREWORKS_API_KEY",
"FIREWORKS_AI_API_KEY",
"FIREWORKSAI_API_KEY",
"FIREWORKS_AI_TOKEN",
"FIREWORKS_API_BASE",
):
monkeypatch.delenv(name, raising=False)
def test_fireworks_ai_provider_config_registration() -> None:
config: Final = ProviderConfigManager.get_provider_responses_api_config(
model="accounts/fireworks/models/kimi-k3", provider=LlmProviders.FIREWORKS_AI
)
assert isinstance(config, FireworksAIResponsesAPIConfig)
assert config.custom_llm_provider == LlmProviders.FIREWORKS_AI
def test_responses_call_hits_native_endpoint_with_mcp_tool_untouched() -> None:
client: Final = _mock_http_client(_fireworks_response("accounts/fireworks/models/kimi-k3"))
mcp_tool: Final[Mcp] = {
"type": "mcp",
"server_label": "deepwiki",
"server_url": "https://mcp.deepwiki.com/mcp",
"require_approval": "never",
}
with patch(HTTPX_CLIENT_FACTORY, return_value=client):
response: Final = litellm.responses(
model="fireworks_ai/accounts/fireworks/models/kimi-k3",
input="What is litellm?",
tools=[mcp_tool], # mutable-ok: the Responses API takes tools as a JSON list
api_key="fw-test-key",
)
url, headers, body = _sent_request(client)
assert url == FIREWORKS_RESPONSES_URL
assert headers["Authorization"] == "Bearer fw-test-key"
assert body["model"] == "accounts/fireworks/models/kimi-k3"
assert tuple(body["tools"]) == (mcp_tool,)
assert "messages" not in body
assert isinstance(response, ResponsesAPIResponse)
function_calls: Final = tuple(item for item in response.output if getattr(item, "type", None) == "function_call")
assert getattr(function_calls[0], "call_id", None) == "call_abc123"
def test_responses_call_expands_bare_model_name_to_fireworks_resource() -> None:
client: Final = _mock_http_client(_fireworks_response("accounts/fireworks/models/glm-5p3"))
with patch(HTTPX_CLIENT_FACTORY, return_value=client):
litellm.responses(model="fireworks_ai/glm-5p3", input="hi", api_key="fw-test-key")
_, _, body = _sent_request(client)
assert body["model"] == "accounts/fireworks/models/glm-5p3"
def test_responses_call_forwards_previous_response_id_and_store() -> None:
client: Final = _mock_http_client(_fireworks_response("accounts/fireworks/models/kimi-k3"))
tool_output: Final[FunctionCallOutput] = {
"type": "function_call_output",
"call_id": "call_abc123",
"output": "{}",
}
with patch(HTTPX_CLIENT_FACTORY, return_value=client):
litellm.responses(
model="fireworks_ai/kimi-k3",
input=[tool_output], # mutable-ok: the Responses API takes input items as a JSON list
previous_response_id="resp_0e946f2d46bf4b49bf8b29ff78083583",
store=True,
api_key="fw-test-key",
)
_, _, body = _sent_request(client)
assert body["previous_response_id"] == "resp_0e946f2d46bf4b49bf8b29ff78083583"
assert body["store"] is True
assert body["input"][0]["call_id"] == "call_abc123"
def test_responses_call_sends_developer_items_as_system_messages() -> None:
client: Final = _mock_http_client(_fireworks_response("accounts/fireworks/models/kimi-k3"))
with patch(HTTPX_CLIENT_FACTORY, return_value=client):
litellm.responses(
model="fireworks_ai/accounts/fireworks/models/kimi-k3",
input=[ # mutable-ok: the Responses API takes input as a JSON list
{"role": "user", "content": "Hi there"},
{"role": "developer", "content": "Answer with exactly one word."},
{"role": "user", "content": [{"type": "input_text", "text": "What is the capital of France?"}]},
],
api_key="fw-test-key",
)
_, _, body = _sent_request(client)
assert tuple(body["input"]) == (
{"role": "user", "content": "Hi there"},
{"role": "system", "content": "Answer with exactly one word.", "type": "message"},
{"role": "user", "content": [{"type": "input_text", "text": "What is the capital of France?"}]},
)
def test_responses_call_maps_pydantic_developer_items_and_replays_pydantic_output_items() -> None:
client: Final = _mock_http_client(_fireworks_response("accounts/fireworks/models/kimi-k3"))
pydantic_input: Final = cast(
ResponseInputParam,
[ # mutable-ok: the Responses API takes input as a JSON list
EasyInputMessage(role="developer", content="Answer with exactly one word.", type="message"),
ResponseReasoningItem(id="rs_1", summary=(), type="reasoning"),
ResponseFunctionToolCall(
id="fc_1",
call_id="call_abc123",
name="get_weather",
arguments='{"city": "Paris"}',
status="completed",
type="function_call",
),
FunctionCallOutput(type="function_call_output", call_id="call_abc123", output="21C"),
],
)
with patch(HTTPX_CLIENT_FACTORY, return_value=client):
litellm.responses(
model="fireworks_ai/accounts/fireworks/models/kimi-k3", input=pydantic_input, api_key="fw-test-key"
)
_, _, body = _sent_request(client)
assert tuple(body["input"]) == (
{"role": "system", "content": "Answer with exactly one word.", "type": "message"},
{"id": "rs_1", "summary": [], "type": "reasoning"},
{
"id": "fc_1",
"call_id": "call_abc123",
"name": "get_weather",
"arguments": '{"city": "Paris"}',
"status": "completed",
"type": "function_call",
},
{"type": "function_call_output", "call_id": "call_abc123", "output": "21C"},
)
def test_file_search_tools_take_litellm_emulated_search_not_fireworks() -> None:
config: Final = FireworksAIResponsesAPIConfig()
file_search: Final = ({"type": "file_search", "vector_store_ids": ("vs_kb",)},)
function_tool: Final = ({"type": "function", "name": "get_weather", "parameters": {"type": "object"}},)
assert should_use_emulated_file_search(tools=file_search, provider_config=config)
assert not should_use_emulated_file_search(tools=function_tool, provider_config=config)
def test_responses_call_sends_session_affinity_for_caller_session_id() -> None:
client: Final = _mock_http_client(_fireworks_response("accounts/fireworks/models/kimi-k3"))
with patch(HTTPX_CLIENT_FACTORY, return_value=client):
litellm.responses(model="fireworks_ai/kimi-k3", input="hi", api_key="fw-test-key", litellm_session_id="sess-42")
_, headers, _ = _sent_request(client)
assert headers["x-session-affinity"] == "sess-42"
def test_responses_call_keeps_caller_supplied_session_affinity_header() -> None:
client: Final = _mock_http_client(_fireworks_response("accounts/fireworks/models/kimi-k3"))
pinned: Final[Mapping[str, str]] = MappingProxyType({"x-session-affinity": "explicit-node"})
with patch(HTTPX_CLIENT_FACTORY, return_value=client):
litellm.responses(
model="fireworks_ai/kimi-k3",
input="hi",
api_key="fw-test-key",
litellm_session_id="sess-42",
extra_headers=pinned,
)
_, headers, _ = _sent_request(client)
assert headers["x-session-affinity"] == "explicit-node"
def test_responses_call_maps_provider_errors_to_fireworks_ai() -> None:
client: Final = MagicMock()
request: Final = httpx.Request("POST", FIREWORKS_RESPONSES_URL)
client.post.side_effect = httpx.HTTPStatusError(
"unauthorized",
request=request,
response=httpx.Response(401, text='{"error": {"message": "invalid api key"}}', request=request),
)
with patch(HTTPX_CLIENT_FACTORY, return_value=client), pytest.raises(litellm.AuthenticationError) as raised:
litellm.responses(model="fireworks_ai/kimi-k3", input="hi", api_key="fw-bad-key")
assert raised.value.llm_provider == "fireworks_ai"
assert raised.value.status_code == 401
assert "invalid api key" in str(raised.value)
def test_responses_call_skips_session_affinity_for_proxy_generated_session_id() -> None:
client: Final = _mock_http_client(_fireworks_response("accounts/fireworks/models/kimi-k3"))
generated: Final[_GeneratedSessionMetadata] = {"litellm_session_id_generated": True}
with patch(HTTPX_CLIENT_FACTORY, return_value=client):
litellm.responses(
model="fireworks_ai/kimi-k3",
input="hi",
api_key="fw-test-key",
litellm_session_id="generated-1",
litellm_metadata=generated,
)
_, headers, _ = _sent_request(client)
assert "x-session-affinity" not in headers
@pytest.mark.parametrize(
"api_base, expected",
(
(None, FIREWORKS_RESPONSES_URL),
("https://api.fireworks.ai/inference/v1", FIREWORKS_RESPONSES_URL),
("https://api.fireworks.ai/inference/v1/", FIREWORKS_RESPONSES_URL),
("https://gateway.example.com/fireworks", "https://gateway.example.com/fireworks/responses"),
),
)
def test_get_complete_url(api_base: str | None, expected: str) -> None:
assert FireworksAIResponsesAPIConfig().get_complete_url(api_base=api_base, litellm_params=NO_PARAMS) == expected
def test_responses_call_reads_fireworks_api_base_env(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("FIREWORKS_API_BASE", "https://self-hosted.example.com/v1")
client: Final = _mock_http_client(_fireworks_response("accounts/fireworks/models/kimi-k3"))
with patch(HTTPX_CLIENT_FACTORY, return_value=client):
litellm.responses(model="fireworks_ai/kimi-k3", input="hi", api_key="fw-test-key")
url, _, _ = _sent_request(client)
assert url == "https://self-hosted.example.com/v1/responses"
@pytest.mark.parametrize(
"env_name", ("FIREWORKS_API_KEY", "FIREWORKS_AI_API_KEY", "FIREWORKSAI_API_KEY", "FIREWORKS_AI_TOKEN")
)
def test_validate_environment_reads_every_fireworks_key_name(monkeypatch: pytest.MonkeyPatch, env_name: str) -> None:
monkeypatch.setenv(env_name, "env-key")
headers: Final = FireworksAIResponsesAPIConfig().validate_environment(
headers=NO_HEADERS, model="accounts/fireworks/models/kimi-k3", litellm_params=GenericLiteLLMParams()
)
assert headers["Authorization"] == "Bearer env-key"
assert headers["Content-Type"] == "application/json"
def test_validate_environment_prefers_explicit_api_key(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("FIREWORKS_API_KEY", "env-key")
headers: Final = FireworksAIResponsesAPIConfig().validate_environment(
headers=NO_HEADERS,
model="accounts/fireworks/models/kimi-k3",
litellm_params=GenericLiteLLMParams(api_key="explicit"),
)
assert headers["Authorization"] == "Bearer explicit"
def test_validate_environment_without_any_key_raises() -> None:
with pytest.raises(ValueError, match="FIREWORKS_API_KEY"):
FireworksAIResponsesAPIConfig().validate_environment(
headers=NO_HEADERS, model="accounts/fireworks/models/kimi-k3", litellm_params=None
)
def test_delete_responses_maps_fireworks_message_only_body_to_deleted_result() -> None:
response_id: Final = (
"resp_xFaIJR9Nc_OXmqKRqL78UuAGj2Te5GY5BT_knpZiMrYoNOVmu5oc2mQW1HI7hCtEYB4mcx2lEYS0DYP1U5yEQskHunuB4=="
)
request: Final = httpx.Request("DELETE", f"{FIREWORKS_RESPONSES_URL}/{quote(response_id, safe='')}")
client: Final = MagicMock()
client.delete.return_value = httpx.Response(200, json={"message": "Response deleted successfully"}, request=request)
with patch(HTTPX_CLIENT_FACTORY, return_value=client):
result: Final = litellm.delete_responses(
response_id=response_id, custom_llm_provider="fireworks_ai", api_key="fw-test-key"
)
assert client.delete.call_args.kwargs["url"] == str(request.url)
assert (result.id, result.object, result.deleted) == (response_id, "response", True)
def _fireworks_stream_response(status: str, output: tuple[Mapping[str, object], ...]) -> Mapping[str, object]:
return {
"id": "resp_htnkJ8piNKeOHkn9LfAusC38O2OgcDQs4S8trSOJ6anLeqjUDGqu2PkWmg5N",
"object": "response",
"created_at": 1788567245,
"model": "accounts/fireworks/models/kimi-k3",
"status": status,
"output": output,
"usage": None
if status == "in_progress"
else {
"input_tokens": 95,
"output_tokens": 89,
"total_tokens": 184,
"input_tokens_details": {"cached_tokens": 94},
},
}
FIREWORKS_SSE_EVENTS: Final[tuple[Mapping[str, object], ...]] = (
{"type": "response.created", "sequence_number": 0, "response": _fireworks_stream_response("in_progress", ())},
{
"type": "response.output_item.added",
"sequence_number": 1,
"output_index": 0,
"item": {"id": "rs_1", "type": "reasoning", "summary": []},
},
{
"type": "response.reasoning_summary_text.delta",
"sequence_number": 2,
"item_id": "rs_1",
"output_index": 0,
"summary_index": 0,
"delta": "pong",
},
{
"type": "response.output_item.added",
"sequence_number": 3,
"output_index": 1,
"item": {"id": "msg_1", "type": "message", "role": "assistant", "status": "in_progress", "content": []},
},
{
"type": "response.output_text.delta",
"sequence_number": 4,
"item_id": "msg_1",
"output_index": 1,
"content_index": 0,
"delta": "po",
},
{
"type": "response.output_text.delta",
"sequence_number": 5,
"item_id": "msg_1",
"output_index": 1,
"content_index": 0,
"delta": "ng",
},
{
"type": "response.completed",
"sequence_number": 6,
"response": _fireworks_stream_response(
"completed",
(
{"id": "rs_1", "type": "reasoning", "summary": [{"type": "summary_text", "text": "pong"}]},
{
"id": "msg_1",
"type": "message",
"role": "assistant",
"status": "completed",
"content": [{"type": "output_text", "text": "pong", "annotations": []}],
},
),
),
},
)
def _sse_body(events: tuple[Mapping[str, object], ...]) -> bytes:
return b"".join(f"data: {json.dumps(dict(event))}\n\n".encode() for event in events) + b"data: [DONE]\n\n"
def test_streaming_responses_call_hits_native_endpoint_and_yields_every_fireworks_event() -> None:
request: Final = httpx.Request("POST", FIREWORKS_RESPONSES_URL)
client: Final = MagicMock()
client.post.return_value = httpx.Response(
200, content=_sse_body(FIREWORKS_SSE_EVENTS), headers={"content-type": "text/event-stream"}, request=request
)
with patch(HTTPX_CLIENT_FACTORY, return_value=client):
received: Final = tuple(
litellm.responses(
model="fireworks_ai/kimi-k3",
input="Reply with the single word pong.",
stream=True,
api_key="fw-test-key",
)
)
url, _, body = _sent_request(client)
assert (url, body["model"], body["stream"], client.post.call_args.kwargs["stream"]) == (
FIREWORKS_RESPONSES_URL,
"accounts/fireworks/models/kimi-k3",
True,
True,
)
assert tuple(event.type for event in received) == tuple(event["type"] for event in FIREWORKS_SSE_EVENTS)
assert "".join(event.delta for event in received if event.type == "response.output_text.delta") == "pong"
assert received[-1].response.usage.output_tokens == 89

View file

@ -12,28 +12,39 @@ import base64
import json
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
from prisma.models import LiteLLM_MCPServerTable as PrismaMCPServer
from litellm.proxy._experimental.mcp_server.db import (
_decode_user_credential,
_prepare_mcp_server_data,
create_mcp_server,
decrypt_credentials,
encrypt_credentials,
get_all_mcp_servers,
get_mcp_servers,
get_mcp_submissions,
get_user_credential,
get_user_oauth_credential,
is_oauth_credential_expired,
list_user_oauth_credentials,
resolve_valid_user_oauth_token,
rotate_mcp_server_credentials_master_key,
rotate_mcp_user_credentials_master_key,
rotate_mcp_user_env_vars_master_key,
store_user_credential,
store_user_oauth_credential,
update_mcp_server,
)
from litellm.proxy._types import NewMCPServerRequest, UpdateMCPServerRequest
from litellm.proxy._types import LiteLLM_MCPServerTable, NewMCPServerRequest, UpdateMCPServerRequest
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
SecretMapDecodeError,
decode_secret_map,
decrypt_value_helper,
encrypt_secret_map,
encrypt_value_helper,
)
from litellm.types.mcp import MCPAuth, MCPTransport
@ -44,6 +55,7 @@ SALT_KEY = "test-salt-key-for-byok-credential-tests-1234"
@pytest.fixture(autouse=True)
def _set_salt_key(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", SALT_KEY)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {"encryption_algorithm": "xsalsa20-poly1305"})
def _make_prisma_with_existing(row):
@ -368,6 +380,169 @@ def test_client_private_key_encrypted_at_rest():
assert decrypted["client_secret"] == "shh"
@pytest.fixture(params=["xsalsa20-poly1305", "aes-256-gcm"])
def map_algorithm(request: pytest.FixtureRequest, monkeypatch: pytest.MonkeyPatch) -> str:
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {"encryption_algorithm": request.param})
return request.param
def _prisma_map_row(data: dict[str, object], quoted: bool = False) -> PrismaMCPServer:
return PrismaMCPServer.model_validate({
"transport": "http", "mcp_access_groups": [], "allowed_tools": [], "extra_headers": [], "args": [],
"allow_all_keys": False, "available_on_public_internet": True, "delegate_auth_to_upstream": False,
"oauth_passthrough": False, "per_server_oauth_discovery": False, "is_byok": False, "byok_description": [],
**data,
**{field: json.dumps(data[field]) for field in ("static_headers", "env") if quoted and data.get(field)},
})
class _MapTable:
def __init__(self, *rows: dict[str, object], quoted: bool = False) -> None:
self.rows = {row["server_id"]: row for row in rows}
self.quoted = quoted
async def create(self, *, data: dict[str, object]) -> PrismaMCPServer:
self.rows = {**self.rows, data["server_id"]: dict(data)}
return _prisma_map_row(data, self.quoted)
async def update(self, *, where: dict[str, str], data: dict[str, object]) -> PrismaMCPServer:
return await self.create(data={**self.rows[where["server_id"]], **data})
async def find_many(self, where: object = None) -> list[PrismaMCPServer]:
return [_prisma_map_row(row, self.quoted) for row in self.rows.values()]
@pytest.mark.asyncio
@pytest.mark.parametrize("field", ["static_headers", "env"])
@pytest.mark.parametrize("quoted", [False, True])
async def test_secret_maps_create_update_round_trip(map_algorithm: str, field: str, quoted: bool) -> None:
table: Final = _MapTable(quoted=quoted)
prisma: Final = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table))
original: Final = {"TOKEN": " sensitive-secret\n", "PREFIX": "v2:gcm:literal", "TEMPLATE": "Bearer ${TOKEN}"}
create: Final = NewMCPServerRequest.model_validate({
"server_id": "srv-map", "transport": "http", "url": "https://up.example.com/mcp", field: original,
})
created: Final = await create_mcp_server(prisma, create, touched_by="test")
first: Final = table.rows["srv-map"][field]
assert isinstance(first, str) and isinstance(json.loads(first), str)
assert json.loads(first).startswith("v2:gcm:") is (map_algorithm == "aes-256-gcm")
assert "sensitive-secret" not in first and "TEMPLATE" not in first
assert getattr(created, field) == original == getattr(create, field)
assert decode_secret_map(first, key=field) == original
replacement: Final = {**original, "TOKEN": "updated-sensitive-secret"}
update: Final = UpdateMCPServerRequest.model_validate({"server_id": "srv-map", field: replacement})
updated: Final = await update_mcp_server(prisma, update, touched_by="test")
second: Final = table.rows["srv-map"][field]
assert second != first and "updated-sensitive-secret" not in second
assert decode_secret_map(second, key=field) == replacement
assert getattr(updated, field) == replacement == getattr(update, field)
assert original["TOKEN"] == " sensitive-secret\n"
omitted: Final = await update_mcp_server(prisma, UpdateMCPServerRequest(server_id="srv-map"), touched_by="test")
assert table.rows["srv-map"][field] == second and getattr(omitted, field) == replacement
cleared: Final = await update_mcp_server(
prisma, UpdateMCPServerRequest.model_validate({"server_id": "srv-map", field: {}}), touched_by="test"
)
assert table.rows["srv-map"][field] == "{}" and getattr(cleared, field) == {}
@pytest.mark.parametrize("field", ["static_headers", "env"])
@pytest.mark.parametrize("as_json", [False, True])
def test_secret_map_legacy_model_read_preserves_exact_values(field: str, as_json: bool) -> None:
original: Final = {"PREFIX": "v2:gcm:literal", "SPACE": " secret\n", "TEMPLATE": "${TOKEN}", "B64": "YWJjZA=="}
incoming: Final = {
"server_id": "srv-map", "transport": "http", field: json.dumps(original) if as_json else original,
}
snapshot: Final = json.dumps(incoming)
parsed: Final = LiteLLM_MCPServerTable.model_validate(incoming)
assert getattr(parsed, field) == original
assert json.dumps(incoming) == snapshot
assert LiteLLM_MCPServerTable.model_validate(parsed.model_dump()).model_dump() == parsed.model_dump()
empty: Final = LiteLLM_MCPServerTable.model_validate({"server_id": "srv-map", "transport": "http", field: None})
assert getattr(empty, field) == ({} if field == "env" else None)
@pytest.mark.parametrize("field", ["static_headers", "env"])
@pytest.mark.parametrize("failure", ["wrong-key", "corrupt", "invalid-values", "invalid-shape", "invalid-json"])
def test_secret_map_model_read_fails_closed(map_algorithm: str, field: str, failure: str) -> None:
plaintext: Final = {"invalid-values": '{"TOKEN": ["sensitive-secret"]}', "invalid-shape": '["sensitive-secret"]',
"invalid-json": "sensitive-secret"}.get(failure, '{"TOKEN": "sensitive-secret"}')
ciphertext: Final = encrypt_value_helper(
plaintext, new_encryption_key="wrong-map-key" if failure == "wrong-key" else None
)
stored: Final = json.dumps(ciphertext[:-8] if failure == "corrupt" else ciphertext)
with pytest.raises(SecretMapDecodeError) as exc:
LiteLLM_MCPServerTable.model_validate({"server_id": "srv-map", "transport": "http", field: stored})
assert field in str(exc.value) and "LITELLM_SALT_KEY" in str(exc.value)
assert all(secret not in str(exc.value) for secret in (plaintext, ciphertext, "sensitive-secret"))
@pytest.mark.asyncio
@pytest.mark.parametrize("field,other", [("static_headers", "env"), ("env", "static_headers")])
async def test_secret_map_rotation_migrates_rekeys_and_preserves_corrupt(
map_algorithm: str, field: str, other: str, monkeypatch: pytest.MonkeyPatch
) -> None:
values: Final = {"TOKEN": "rotation-sensitive-secret", "TEMPLATE": "Bearer ${TOKEN}"}
old: Final = encrypt_secret_map(values)
corrupt: Final = json.dumps(json.loads(old)[:-8])
table: Final = _MapTable(
{"server_id": "broken", field: corrupt, other: old},
{"server_id": "legacy", field: json.dumps(values), other: "{}"},
{"server_id": "encrypted", field: old, other: None},
)
prisma: Final = SimpleNamespace(db=SimpleNamespace(
litellm_mcpservertable=table, litellm_mcpserveroauthclient=SimpleNamespace(find_many=AsyncMock(return_value=[]))
))
await rotate_mcp_server_credentials_master_key(prisma, touched_by="test", new_master_key="rotated-map-key")
assert table.rows["broken"][field] == corrupt
assert table.rows["legacy"][other] == "{}" and table.rows["encrypted"][other] is None
for server_id, map_field in (("broken", other), ("legacy", field), ("encrypted", field)):
stored: Final = table.rows[server_id][map_field]
assert isinstance(json.loads(stored), str) and stored != old and "rotation-sensitive-secret" not in stored
with pytest.raises(SecretMapDecodeError):
decode_secret_map(stored, key=map_field)
monkeypatch.setenv("LITELLM_SALT_KEY", "rotated-map-key")
for server_id, map_field in (("broken", other), ("legacy", field), ("encrypted", field)):
assert decode_secret_map(table.rows[server_id][map_field], key=map_field) == values
@pytest.mark.asyncio
@pytest.mark.parametrize("reader", [get_all_mcp_servers, get_mcp_servers, get_mcp_submissions])
@pytest.mark.parametrize("field", ["static_headers", "env"])
async def test_bulk_reads_isolate_corrupt_secret_maps(reader, field, map_algorithm, caplog):
secret = {"TOKEN": "bulk-sensitive-secret"}
encrypted = encrypt_secret_map(secret)
corrupt = encrypt_secret_map(secret, new_encryption_key="wrong-bulk-key")
rows = [
_prisma_map_row({"server_id": "broken", field: corrupt, "approval_status": "pending_review"}),
_prisma_map_row({"server_id": "healthy", field: encrypted, "approval_status": "active"}),
]
snapshot = [row.model_dump() for row in rows]
table = SimpleNamespace(find_many=AsyncMock(return_value=rows))
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table))
result = await reader(prisma, ["broken", "healthy"]) if reader is get_mcp_servers else await reader(prisma)
items = result.items if reader is get_mcp_submissions else result
assert [row.server_id for row in items] == ["healthy"]
assert getattr(items[0], field) == secret
assert [row.model_dump() for row in rows] == snapshot
assert "broken" in caplog.text
assert all(value not in caplog.text for value in ("bulk-sensitive-secret", corrupt, encrypted))
if reader is get_mcp_submissions:
assert (result.total, result.pending_review, result.active, result.rejected) == (1, 0, 1, 0)
@pytest.mark.asyncio
@pytest.mark.parametrize("reader", [get_all_mcp_servers, get_mcp_servers, get_mcp_submissions])
async def test_bulk_reads_do_not_swallow_unrelated_validation_errors(reader):
from pydantic import ValidationError
row = _prisma_map_row({"server_id": "invalid", "transport": "unsupported"})
table = SimpleNamespace(find_many=AsyncMock(return_value=[row]))
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table))
request = reader(prisma, ["invalid"]) if reader is get_mcp_servers else reader(prisma)
with pytest.raises(ValidationError, match="transport"):
await request
# ── BYOK round-trip ───────────────────────────────────────────────────────────

View file

@ -861,6 +861,7 @@ _SALT_KEY = "test-salt-key-for-env-vars-tests-1234"
@pytest.fixture
def env_vars_salt_key(monkeypatch):
monkeypatch.setenv("LITELLM_SALT_KEY", _SALT_KEY)
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {"encryption_algorithm": "xsalsa20-poly1305"})
def _mock_env_vars_prisma(row=None):
@ -1518,9 +1519,16 @@ async def test_create_mcp_server_decrypts_env_vars_when_prisma_returns_json_stri
assert "s3cr3t-p@ss" not in encrypted_env_vars_str
def _prisma_row_with_json_string_env_vars():
row = MagicMock()
row.env_vars = encrypted_env_vars_str
return row
import json
from prisma.models import LiteLLM_MCPServerTable
return LiteLLM_MCPServerTable.model_validate({
"server_id": "srv-returned", "transport": "http", "mcp_access_groups": [], "allowed_tools": [],
"extra_headers": [], "args": [], "allow_all_keys": False, "available_on_public_internet": True,
"delegate_auth_to_upstream": False, "oauth_passthrough": False, "per_server_oauth_discovery": False,
"is_byok": False, "byok_description": [], "env_vars": json.dumps(encrypted_env_vars_str),
})
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.create = AsyncMock(
@ -1537,7 +1545,9 @@ async def test_create_mcp_server_decrypts_env_vars_when_prisma_returns_json_stri
touched_by="test-user",
)
assert isinstance(created.env_vars, list)
assert created.env_vars[0]["value"] == "s3cr3t-p@ss"
assert created.env_vars[0].value == "s3cr3t-p@ss"
assert created.env_vars[0].name == "DB_PASSWORD"
assert created.env == {}
mock_prisma_upd = MagicMock()
mock_prisma_upd.db.litellm_mcpservertable.update = AsyncMock(
@ -1549,7 +1559,9 @@ async def test_create_mcp_server_decrypts_env_vars_when_prisma_returns_json_stri
touched_by="test-user",
)
assert isinstance(updated.env_vars, list)
assert updated.env_vars[0]["value"] == "s3cr3t-p@ss"
assert updated.env_vars[0].value == "s3cr3t-p@ss"
assert updated.env_vars[0].name == "DB_PASSWORD"
assert updated.env == {}
def test_reencrypt_global_env_var_values_handles_json_string(env_vars_salt_key):

View file

@ -11,7 +11,7 @@ import json
from unittest.mock import AsyncMock, MagicMock
import pytest
from prisma import Json
from prisma import Json, models
from litellm.proxy._experimental.mcp_server.db import (
create_mcp_server,
@ -28,8 +28,11 @@ def _credentials_cleared(value) -> bool:
def _mock_prisma():
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable = AsyncMock()
mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=MagicMock())
mock_prisma.db.litellm_mcpservertable.create = AsyncMock(return_value=MagicMock())
row = models.LiteLLM_MCPServerTable.model_construct(
server_id="test-server", transport="http", env={}, env_vars=[]
)
mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=row)
mock_prisma.db.litellm_mcpservertable.create = AsyncMock(return_value=row)
return mock_prisma

View file

@ -905,6 +905,68 @@ class TestMCPServerManager:
assert retry_slot is not None
assert retry_slot.generation > old_generation
@pytest.mark.asyncio
@pytest.mark.parametrize("corrupt_column", ("static_headers", "env"))
async def test_database_reload_drops_cached_server_whose_secret_map_stops_decoding(
self, monkeypatch, caplog, corrupt_column
):
from types import SimpleNamespace
from litellm.proxy import proxy_server
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_secret_map
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-reload-secret-map-salt")
monkeypatch.setattr(proxy_server, "general_settings", {"encryption_algorithm": "aes-256-gcm"})
headers = {"Authorization": "Bearer dummy-header-secret-4f1c"}
env = {"UPSTREAM_TOKEN": "dummy-env-secret-9a2b"}
stamp = datetime.now()
cached = MCPServer(
server_id="cached-server",
name="cached_server",
url="https://up.example.com/mcp",
transport=MCPTransport.http,
static_headers=dict(headers),
env=dict(env),
updated_at=stamp,
)
manager = MCPServerManager()
manager.registry[cached.server_id] = cached
stored = {"static_headers": encrypt_secret_map(headers), "env": encrypt_secret_map(env)}
corrupted = {**stored, corrupt_column: stored[corrupt_column][:-6] + 'AAAAA"'}
def _row(server_id, maps):
row = MagicMock()
row.server_id = server_id
row.alias = server_id
row.model_dump.return_value = {
"server_id": server_id,
"alias": server_id,
"server_name": server_id,
"url": "https://up.example.com/mcp",
"transport": MCPTransport.http,
"updated_at": stamp,
**maps,
}
return row
table = SimpleNamespace(
find_many=AsyncMock(return_value=[_row(cached.server_id, corrupted), _row("healthy-sibling", stored)])
)
prisma = SimpleNamespace(db=SimpleNamespace(litellm_mcpservertable=table))
monkeypatch.setattr(proxy_server, "prisma_client", prisma)
with caplog.at_level(logging.DEBUG, logger="LiteLLM"):
await manager.reload_servers_from_database()
assert set(manager.registry) == {"healthy-sibling"}
sibling = manager.registry["healthy-sibling"]
assert dict(sibling.static_headers) == headers
assert dict(sibling.env) == env
logged = "\n".join(caplog.messages)
assert cached.server_id in logged
for secret in (*headers.values(), *env.values(), *stored.values(), corrupted[corrupt_column]):
assert secret not in logged
@pytest.mark.asyncio
async def test_lazy_oauth_discovery_preserves_manual_authorization_url_gate(self):
with patch.dict(os.environ, {"LITELLM_MCP_OAUTH_DISCOVERY_ON_STARTUP": "false"}):

View file

@ -15,6 +15,11 @@ import httpx
from litellm.experimental_mcp_client.client import MCPSigV4Auth, MCPClient
from litellm.types.mcp import MCPAuth, MCPTransport
from prisma import models
def _updated_row() -> models.LiteLLM_MCPServerTable:
return models.LiteLLM_MCPServerTable.model_construct(server_id="test-server", transport="http", env={}, env_vars=[])
class TestMCPSigV4Auth:
@ -600,7 +605,7 @@ class TestCredentialMergeOnUpdate:
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing_record)
mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=MagicMock())
mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=_updated_row())
data = UpdateMCPServerRequest(
server_id="test-server",
@ -639,7 +644,7 @@ class TestCredentialMergeOnUpdate:
from litellm.proxy._types import UpdateMCPServerRequest
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=MagicMock())
mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=_updated_row())
data = UpdateMCPServerRequest(
server_id="test-server",
@ -667,7 +672,7 @@ class TestCredentialMergeOnUpdate:
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing_record)
mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=MagicMock())
mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=_updated_row())
data = UpdateMCPServerRequest(
server_id="test-server",
@ -709,7 +714,7 @@ class TestCredentialMergeOnUpdate:
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing_record)
mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=MagicMock())
mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=_updated_row())
data = UpdateMCPServerRequest(
server_id="test-server",
@ -752,7 +757,7 @@ class TestCredentialMergeOnUpdate:
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing_record)
mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=MagicMock())
mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=_updated_row())
data = UpdateMCPServerRequest(
server_id="test-server",
@ -1083,7 +1088,7 @@ class TestAuthTypeSwitchClearsCredentials:
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_unique = AsyncMock(return_value=existing_record)
mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=MagicMock())
mock_prisma.db.litellm_mcpservertable.update = AsyncMock(return_value=_updated_row())
data = UpdateMCPServerRequest(
server_id="test-server",

View file

@ -8,6 +8,7 @@ proof-of-fix (real proxy + DB) is performed separately on the repro server.
import json
from types import SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -457,6 +458,65 @@ async def test_scan_covered_tables_classifies_legacy_and_v2(salt_key, monkeypatc
assert by_loc["credentials"].legacy == 0
@pytest.mark.asyncio
@pytest.mark.parametrize("column", ("static_headers", "env"))
@pytest.mark.parametrize("algorithm", ("xsalsa20-poly1305", "aes-256-gcm"))
@pytest.mark.parametrize("as_json", (False, True))
@pytest.mark.parametrize(
"case", ("legacy", "encrypted", "wrong-key", "corrupt", "invalid-shape", "invalid-scalar", "empty", "null")
)
async def test_check_classifies_mcp_secret_maps(
salt_key: str,
monkeypatch: pytest.MonkeyPatch,
column: str,
algorithm: str,
as_json: bool,
case: str,
) -> None:
monkeypatch.setattr(proxy_server, "general_settings", {"encryption_algorithm": algorithm})
plaintext: Final = {"Authorization": "v2:gcm:operator-text", "CUSTOM": "litellm_enc::literal\n café "}
ciphertext: Final = encrypt_value_helper(json.dumps(plaintext))
cases: Final[dict[str, object]] = {
"legacy": plaintext,
"encrypted": ciphertext,
"wrong-key": encrypt_value_helper(json.dumps(plaintext), new_encryption_key="different-map-salt"),
"corrupt": ciphertext[:-4] + "AAAA",
"invalid-shape": encrypt_value_helper(json.dumps({"Authorization": 42})),
"invalid-scalar": "null",
"empty": {},
"null": None,
}
value: Final = json.dumps(cases[case]) if as_json and case != "null" else cases[case]
row: Final = SimpleNamespace(**{column: value})
client: Final = MagicMock()
_empty_covered_tables(client)
client.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[row])
client.db.litellm_mcpservertable.update = AsyncMock()
client.db.litellm_teamtable.find_many = AsyncMock(return_value=[])
client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None)
client.db.litellm_config.find_unique = AsyncMock(return_value=None)
monkeypatch.setattr(proxy_server, "general_settings", {})
report: Final = await cm.check_encryption(client)
expected_legacy: Final = int(case == "legacy" or (case == "encrypted" and algorithm == "xsalsa20-poly1305"))
expected_v2: Final = int(case == "encrypted" and algorithm == "aes-256-gcm")
expected_invalid: Final = int(case in ("wrong-key", "corrupt", "invalid-shape", "invalid-scalar"))
assert report.as_dict()["locations"]["mcp_server"] == {
"scanned": int(case not in ("empty", "null")),
"migrated": 0,
"already_v2": expected_v2,
"plaintext": 0,
"undecryptable": expected_invalid,
"legacy": expected_legacy,
}
assert report.residual_legacy == expected_legacy
assert report.total_undecryptable == expected_invalid
assert getattr(row, column) == value
client.db.litellm_mcpservertable.update.assert_not_awaited()
@pytest.mark.asyncio
async def test_check_counts_covered_table_residual(salt_key, monkeypatch):
"""check_encryption now scans the rotation-covered tables (model table here),

View file

@ -17,6 +17,9 @@ from litellm.llms.chatgpt.responses.transformation import ChatGPTResponsesAPICon
from litellm.llms.databricks.responses.transformation import (
DatabricksResponsesAPIConfig,
)
from litellm.llms.fireworks_ai.responses.transformation import (
FireworksAIResponsesAPIConfig,
)
from litellm.llms.github_copilot.responses.transformation import (
GithubCopilotResponsesAPIConfig,
)
@ -102,6 +105,12 @@ class TestResponsesAPIWebSocketSupport:
def test_openai_model_in_websocket_url_default(self):
assert OpenAIResponsesAPIConfig().model_in_websocket_url() is True
def test_fireworks_ai_uses_managed_websocket(self):
"""Fireworks AI should use managed websocket handler"""
assert (
FireworksAIResponsesAPIConfig().supports_native_websocket() is False
), "Fireworks AI should use managed websocket handler"
def test_xai_uses_managed_websocket(self):
"""XAI should use managed websocket handler"""
config = XAIResponsesAPIConfig()