mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
Merge branch 'litellm_internal_staging' into litellm_lit_4116_drop_params_string_coerce
This commit is contained in:
commit
b827375e60
18 changed files with 1058 additions and 69 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
102
litellm/llms/fireworks_ai/responses/transformation.py
Normal file
102
litellm/llms/fireworks_ai/responses/transformation.py
Normal 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
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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 ───────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"}):
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue