Merge remote-tracking branch 'origin/litellm_internal_staging' into feat-prism-provider

This commit is contained in:
rajitkhanna 2026-09-11 19:14:54 -07:00
commit 549ef1ebe5
14 changed files with 543 additions and 9 deletions

View file

@ -257,6 +257,14 @@ IAM_TOKEN_DB_AUTH / AZURE_POSTGRESQL_AUTH toggle that only the writer sets.
- name: DATABASE_SCHEMA
value: {{ .schema | quote }}
{{- end }}
{{- if .sslMode }}
- name: DATABASE_SSLMODE
value: {{ .sslMode | quote }}
{{- end }}
{{- if .sslRootCert }}
- name: DATABASE_SSLROOTCERT
value: {{ .sslRootCert | quote }}
{{- end }}
{{- if and .useIAMAuth .useAzureEntraAuth }}
{{- fail "database.writer.useIAMAuth and database.writer.useAzureEntraAuth are mutually exclusive: the database password can only come from one token source" }}
{{- end }}

View file

@ -4,6 +4,7 @@ templates:
- gateway/configmap.yaml
- backend/deployment.yaml
- backend/configmap.yaml
- migrations-job.yaml
values:
- ./values/required.yaml
tests:
@ -67,6 +68,82 @@ tests:
value: "true"
any: true
- it: emits no TLS env by default
template: gateway/deployment.yaml
asserts:
- notContains:
path: spec.template.spec.containers[0].env
content:
name: DATABASE_SSLMODE
any: true
- notContains:
path: spec.template.spec.containers[0].env
content:
name: DATABASE_SSLROOTCERT
any: true
- it: writer sslMode and sslRootCert reach gateway and backend as DATABASE_SSLMODE and DATABASE_SSLROOTCERT
templates:
- gateway/deployment.yaml
- backend/deployment.yaml
set:
database.writer.useIAMAuth: true
database.writer.sslMode: verify-full
database.writer.sslRootCert: /etc/ssl/certs/ca-certificates.crt
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: DATABASE_SSLMODE
value: verify-full
any: true
- contains:
path: spec.template.spec.containers[0].env
content:
name: DATABASE_SSLROOTCERT
value: /etc/ssl/certs/ca-certificates.crt
any: true
- it: writer sslMode and sslRootCert reach the collector sidecar and the migrations job, which dial Postgres themselves
set:
gateway.collector.enabled: true
database.connectionPool.enabled: true
database.writer.sslMode: verify-full
database.writer.sslRootCert: /etc/ssl/certs/ca-certificates.crt
asserts:
- equal:
path: spec.template.spec.containers[1].name
value: collector
template: gateway/deployment.yaml
- contains:
path: spec.template.spec.containers[1].env
content:
name: DATABASE_SSLMODE
value: verify-full
any: true
template: gateway/deployment.yaml
- contains:
path: spec.template.spec.containers[1].env
content:
name: DATABASE_SSLROOTCERT
value: /etc/ssl/certs/ca-certificates.crt
any: true
template: gateway/deployment.yaml
- contains:
path: spec.template.spec.containers[0].env
content:
name: DATABASE_SSLMODE
value: verify-full
any: true
template: migrations-job.yaml
- contains:
path: spec.template.spec.containers[0].env
content:
name: DATABASE_SSLROOTCERT
value: /etc/ssl/certs/ca-certificates.crt
any: true
template: migrations-job.yaml
- it: writer rejects both token sources at once
template: gateway/deployment.yaml
set:

View file

@ -208,6 +208,11 @@ database:
name: litellm-writer-secret
usernameKey: username
passwordKey: password
# libpq sslmode / sslrootcert applied to the writer and reader URLs (Prisma and the
# in-container PgBouncer); e.g. verify-full with /etc/ssl/certs/ca-certificates.crt for AWS RDS.
# sslRootCert on its own implies sslMode verify-full
sslMode: ""
sslRootCert: ""
# Optional read-replica routing. When `reader.host` is set, the proxy routes
# reads (find_*, count, group_by, query_raw/_first) to this endpoint while

View file

@ -27,6 +27,7 @@ from litellm._logging import verbose_logger
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import INTERNAL_CALL_ORIGIN_METADATA_KEY
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.websearch_interception.tools import is_web_search_tool_responses
from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs
from litellm.litellm_core_utils.internal_call_metadata import sanitized_forwardable_call_metadata
from litellm.litellm_core_utils.llm_judge import (
@ -335,6 +336,16 @@ def _forwards_nothing(value: object) -> bool:
return value is None or (isinstance(value, list) and len(value) == 0)
def _request_has_hosted_web_search(request: Mapping[str, object]) -> bool:
if request.get("web_search_options") is not None:
return True
tools: Final = request.get("tools")
return isinstance(tools, Sequence) and any(
isinstance(tool, Mapping) and tool.get("type") != "function" and is_web_search_tool_responses(tool)
for tool in tools
)
def _judgeable_sample(
ops: _SurfaceOps,
kwargs: Mapping[str, object],
@ -343,9 +354,14 @@ def _judgeable_sample(
) -> tuple[tuple[Mapping[str, object], ...], Mapping[str, object], str] | None:
"""The normalized chat conversation, the forwardable generation params, and the
judgeable final text; None when this request's shapes cannot be sampled (no text and no
tool call to serialize, or a shape the owner transformations reject)."""
tool call to serialize, hosted web search the shadow cannot replay comparably,
or a shape the owner transformations reject)."""
if _request_has_hosted_web_search(_proxy_wire_body(kwargs) if ops.wire_params else model_parameters):
return None
try:
request: Final = ops.chat_request(kwargs, model_parameters)
if _request_has_hosted_web_search(request):
return None
items: Final = _MESSAGE_ITEMS_ADAPTER.validate_python(request.get("messages"))
messages: Final = _CHAT_MESSAGES_ADAPTER.validate_python(
tuple(m.model_dump(exclude_none=True) if isinstance(m, BaseModel) else m for m in items)

View file

@ -400,7 +400,7 @@ class LiteLLMAnthropicMessagesAdapter:
Anthropic web search tools have:
- type starting with "web_search" (e.g., "web_search_20260209")
- name = "web_search"
- legacy name = "web_search" without a client input_schema
Args:
tool: Tool definition dict
@ -410,7 +410,9 @@ class LiteLLMAnthropicMessagesAdapter:
"""
tool_type: Final = tool.get("type", "")
tool_name: Final = tool.get("name", "")
return (isinstance(tool_type, str) and tool_type.startswith("web_search")) or tool_name == "web_search"
return (isinstance(tool_type, str) and tool_type.startswith("web_search")) or (
tool_name == "web_search" and "input_schema" not in tool
)
def translate_anthropic_messages_to_openai(
self,

View file

@ -74,6 +74,8 @@ DisablePreparedStatementsFlag = Annotated[
bool, BeforeValidator(partial(token_auth_flag_enabled, env_var=DISABLE_PREPARED_STATEMENTS_ENV_VAR))
]
MAX_IDLE_CONNECTION_LIFETIME_ENV_VAR: Final = "DATABASE_MAX_IDLE_CONNECTION_LIFETIME"
DATABASE_SSLMODE_ENV_VAR: Final = "DATABASE_SSLMODE"
DATABASE_SSLROOTCERT_ENV_VAR: Final = "DATABASE_SSLROOTCERT"
# schema.prisma pins `provider = "postgresql"`, so these are the only schemes
# Prisma can actually connect with.
@ -135,6 +137,7 @@ def add_missing_query_params(url: str, params: Mapping[str, str | int | float])
LIBPQ_VERIFY_SSLMODES: Final[frozenset[str]] = frozenset({"verify-ca", "verify-full"})
PRISMA_TLS_PARAM_KEYS: Final[frozenset[str]] = frozenset({"sslmode", "sslcert", "sslaccept"})
PEM_CERT_HEADER: Final = b"-----BEGIN CERTIFICATE-----"
PG_SSL_REQUEST: Final = struct.pack("!ii", 8, 80877103)
TLS_PROBE_TIMEOUT_SECONDS: Final = 10.0
@ -263,6 +266,19 @@ def connection_params_from_url(url: str) -> Mapping[str, str | int | float]:
)
def token_refresh_params_from_url(url: str) -> Mapping[str, str | int | float]:
"""Return the params a re-minted token URL carries over from the URL it replaces.
The pool and timeout params plus Prisma's TLS params (already translated from
libpq spelling), so a refreshed URL keeps verifying the server the way the
first one did.
"""
kept: Final = CONNECTION_PARAM_KEYS | PRISMA_TLS_PARAM_KEYS
return MappingProxyType(
{key: value for key, value in urllib.parse.parse_qsl(urllib.parse.urlsplit(url).query) if key in kept}
)
def unsupported_db_scheme(database_url: str) -> str | None:
"""Return the connection URL scheme when it is not PostgreSQL, else None.
@ -312,6 +328,9 @@ class DatabaseURLSettings(BaseSettings):
default=None, validation_alias=MAX_IDLE_CONNECTION_LIFETIME_ENV_VAR
)
database_sslmode: str | None = Field(default=None, validation_alias=DATABASE_SSLMODE_ENV_VAR)
database_sslrootcert: str | None = Field(default=None, validation_alias=DATABASE_SSLROOTCERT_ENV_VAR)
# Writer
database_url: str | None = Field(default=None, validation_alias="DATABASE_URL")
direct_url: str | None = Field(default=None, validation_alias="DIRECT_URL")
@ -353,6 +372,25 @@ class DatabaseURLSettings(BaseSettings):
azure_postgresql_auth=self.azure_postgresql_auth,
)
def tls_params(self) -> Mapping[str, str]:
"""``sslmode`` / ``sslrootcert`` query params for every URL assembled from the discrete vars.
A root cert on its own means ``verify-full``: under libpq's default
``prefer`` the CA would never be consulted, and PgBouncer would dial
Postgres unverified with the bundle loaded.
"""
sslmode: Final = self.database_sslmode or ("verify-full" if self.database_sslrootcert else None)
return MappingProxyType(
{
key: value
for key, value in (
("sslmode", sslmode),
("sslrootcert", self.database_sslrootcert),
)
if value
}
)
def build_writer_url(self) -> str | None:
"""Return the writer URL to set, or ``None`` to leave it as-is.
@ -362,6 +400,12 @@ class DatabaseURLSettings(BaseSettings):
A ``DATABASE_URL`` the supervisor pointed at the in-container PgBouncer
is kept even under token auth: the pooler renews the token upstream.
"""
assembled: Final = self._assemble_writer_url()
if assembled is None:
return None
return add_missing_query_params(assembled, self.tls_params())
def _assemble_writer_url(self) -> str | None:
auth: Final = self.token_auth()
if auth is not None and database_url_is_pooled():
return None
@ -411,6 +455,12 @@ class DatabaseURLSettings(BaseSettings):
pre-existing ``DATABASE_URL_READ_REPLICA``. Reader fields fall back
to the writer's values.
"""
assembled: Final = self._assemble_reader_url()
if assembled is None:
return None
return add_missing_query_params(assembled, self.tls_params())
def _assemble_reader_url(self) -> str | None:
if not self.database_host_read_replica:
return None # reader is opt-in
if self.database_url_read_replica:

View file

@ -16,7 +16,7 @@ from datetime import datetime, timedelta
from typing import Any, Final, Protocol
from litellm._logging import verbose_proxy_logger
from litellm.proxy.db.db_url_settings import add_missing_query_params, connection_params_from_url
from litellm.proxy.db.db_url_settings import add_missing_query_params, token_refresh_params_from_url
from litellm.proxy.db.token_auth import (
DEFAULT_POSTGRES_PORT,
DatabaseTokenAuth,
@ -441,7 +441,7 @@ class PrismaWrapper:
endpoint: Final = self._iam_endpoint if self._iam_endpoint is not None else self._endpoint_from_env()
db_url: Final = add_missing_query_params(
endpoint.build_url(mint_database_token(auth, endpoint)),
connection_params_from_url(os.environ.get(self._db_url_env_var, "")),
token_refresh_params_from_url(os.environ.get(self._db_url_env_var, "")),
)
os.environ[self._db_url_env_var] = db_url
return db_url

View file

@ -121,6 +121,11 @@ from litellm.proxy.db.create_views import (
should_create_missing_views,
)
from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter
from litellm.proxy.db.db_url_settings import (
DatabaseURLSettings,
add_missing_query_params,
token_refresh_params_from_url,
)
from litellm.proxy.db.exception_handler import (
PrismaDBExceptionHandler,
call_with_db_reconnect_retry,
@ -4054,7 +4059,10 @@ class PrismaClient:
# loop and times out after 30s.
if token_auth is not None and reader_iam_endpoint is not None:
reader_token: Final = mint_database_token(token_auth, reader_iam_endpoint)
read_replica_url = reader_iam_endpoint.build_url(reader_token)
read_replica_url = add_missing_query_params(
reader_iam_endpoint.build_url(reader_token),
token_refresh_params_from_url(read_replica_url),
)
os.environ["DATABASE_URL_READ_REPLICA"] = read_replica_url
reader_kwargs: Final[dict[str, Any]] = {"datasource": {"url": read_replica_url}}
if http_client is not None:
@ -7807,7 +7815,7 @@ def construct_database_url_from_env_vars() -> str | None:
if database_schema:
database_url += f"?schema={database_schema}"
return database_url
return add_missing_query_params(database_url, DatabaseURLSettings.from_env().tls_params())
return None

View file

@ -2,6 +2,7 @@
the detached pipeline's single attempt-row write, and the cache-first job lookup."""
import asyncio
from collections.abc import Mapping
from datetime import datetime, timedelta, timezone
from typing import Final
from unittest.mock import AsyncMock, MagicMock
@ -421,6 +422,183 @@ class TestSurfaceNormalization:
assert "previous_response_id" not in shadow_call
assert "instructions" not in shadow_call
@pytest.mark.parametrize(
"call_type,search_params,model",
[
("completion", {"web_search_options": {}}, "anthropic/claude-fable-5"),
(
"acompletion",
{"web_search_options": {"search_context_size": "high"}},
"anthropic/claude-fable-5",
),
(
"acompletion",
{"tools": [{"type": "web_search_20260209", "name": "web_search"}]},
"anthropic/claude-fable-5",
),
(
"anthropic_messages",
{"tools": [{"type": "web_search_20250305", "name": "web_search"}]},
"anthropic/claude-fable-5",
),
(
"anthropic_messages",
{"tools": [{"type": "web_search_20260209", "name": "web_search"}]},
"anthropic/claude-fable-5",
),
(
"anthropic_messages",
{"tools": [{"name": "web_search"}]},
"anthropic/claude-fable-5",
),
("aresponses", {"tools": [{"type": "web_search"}]}, "anthropic/claude-fable-5"),
("responses", {"tools": [{"type": "web_search_preview"}]}, "anthropic/claude-fable-5"),
("aresponses", {"tools": [{"type": "web_search_2025_08_26"}]}, "anthropic/claude-fable-5"),
(
"responses",
{"tools": [{"type": "web_search_preview_2025_03_11"}]},
"anthropic/claude-fable-5",
),
("aresponses", {"tools": [{"type": "web_search"}]}, "bedrock/us.anthropic.claude-fable-5"),
(
"responses",
{"tools": [{"type": "web_search_preview"}]},
"bedrock/us.anthropic.claude-fable-5",
),
(
"acompletion",
{
"tools": [
{"type": "function", "function": {"name": "WebSearch", "parameters": {"type": "object"}}},
{"type": "web_search_20260209", "name": "web_search"},
]
},
"anthropic/claude-fable-5",
),
],
ids=[
"chat-empty-options",
"chat-configured-options",
"chat-provider-transformed-tools",
"messages-native-search",
"messages-dated-search",
"messages-legacy-search-normalized",
"responses-search",
"responses-preview",
"responses-dated-search",
"responses-dated-preview",
"responses-bedrock-erases-search",
"responses-bedrock-erases-preview",
"chat-mixed-client-and-hosted-tools",
],
)
async def test_hosted_web_search_skips_shadow_calls_and_spend(
self, call_type: str, search_params: Mapping[str, object], model: str
) -> None:
base_kwargs: Final = _success_kwargs(call_type=call_type, model=model)
is_chat: Final = call_type in ("completion", "acompletion")
is_responses: Final = call_type in ("responses", "aresponses")
hook_kwargs: Final = {
**base_kwargs,
"model": model,
"messages": "what is new" if is_responses else base_kwargs["messages"],
"standard_logging_object": {
**base_kwargs["standard_logging_object"],
"model_parameters": search_params if is_chat else {},
},
"litellm_params": {
**base_kwargs["litellm_params"],
"proxy_server_request": {"body": {} if is_chat else search_params},
},
}
prisma: Final = _prisma()
router: Final = _router()
counter: Final = {"spend:shadow_eval:job-1": 0.1, "spend:shadow_eval:job-2": 0.1}
logger: Final = _logger(
router=router,
prisma=prisma,
jobs=(_job(max_budget=0.2), _job(id="job-2", max_budget=0.2)),
counter_store=counter,
)
await logger.async_log_success_event(
hook_kwargs, RESPONSES_API_RESPONSE if is_responses else RESPONSE, None, None
)
await _drain(logger)
router.acompletion.assert_not_called()
prisma.db.litellm_shadowevalattempt.create.assert_not_called()
assert logger._test_funnel == [("job-1", "unjudgeable"), ("job-2", "unjudgeable")]
assert logger._job_starts == {}
assert logger._test_counter == {"spend:shadow_eval:job-1": 0.1, "spend:shadow_eval:job-2": 0.1}
@pytest.mark.parametrize(
"call_type,tool_name",
[
(call_type, tool_name)
for call_type in ("completion", "acompletion", "anthropic_messages", "responses", "aresponses")
for tool_name in ("WebSearch", "litellm_web_search", "web_search")
],
)
async def test_client_web_search_tools_remain_sampled(self, call_type: str, tool_name: str) -> None:
is_chat: Final = call_type in ("completion", "acompletion")
is_responses: Final = call_type in ("responses", "aresponses")
tool: Final = (
{"type": "function", "function": {"name": tool_name, "parameters": {"type": "object"}}}
if is_chat
else {"type": "function", "name": tool_name, "parameters": {"type": "object"}}
if is_responses
else {"name": tool_name, "input_schema": {"type": "object", "properties": {}}}
)
source: Final = {"tools": [tool], "web_search_options": None}
base_kwargs: Final = _success_kwargs(call_type=call_type)
hook_kwargs: Final = {
**base_kwargs,
"messages": "search for current news" if is_responses else base_kwargs["messages"],
"standard_logging_object": {
**base_kwargs["standard_logging_object"],
"model_parameters": source if is_chat else {},
},
"litellm_params": {
**base_kwargs["litellm_params"],
"proxy_server_request": {"body": {} if is_chat else source},
},
}
prisma, router = await self._drive(hook_kwargs, RESPONSES_API_RESPONSE if is_responses else RESPONSE)
assert router.acompletion.call_count == 2
shadow_call: Final = router.acompletion.call_args_list[0].kwargs
assert shadow_call["tools"][0]["function"]["name"] == tool_name
assert "web_search_options" not in shadow_call
prisma.db.litellm_shadowevalattempt.create.assert_called_once()
@pytest.mark.parametrize("call_type", ["completion", "acompletion"])
async def test_chat_search_removed_by_guardrail_still_samples(self, call_type: str) -> None:
base_kwargs: Final = _success_kwargs(
call_type=call_type,
request_metadata={
"standard_logging_guardrail_information": [{"guardrail_name": "g", "guardrail_mode": "pre_call"}]
},
)
hook_kwargs: Final = {
**base_kwargs,
"litellm_params": {
**base_kwargs["litellm_params"],
"proxy_server_request": {
"body": {"web_search_options": {}, "tools": [{"type": "web_search_20260209"}]}
},
},
}
prisma, router = await self._drive(hook_kwargs, RESPONSE)
shadow_call: Final = router.acompletion.call_args_list[0].kwargs
assert "web_search_options" not in shadow_call
assert "tools" not in shadow_call
assert router.acompletion.call_count == 2
prisma.db.litellm_shadowevalattempt.create.assert_called_once()
@pytest.mark.parametrize("payload_shape", ["typed", "dict"])
@pytest.mark.parametrize("call_type", ["aresponses", "responses"])
async def test_responses_arms_normalize_bare_string_input_and_instructions(self, call_type, payload_shape):

View file

@ -3358,6 +3358,27 @@ def test_is_web_search_tool():
assert adapter._is_web_search_tool(regular_tool) is False
@pytest.mark.parametrize("schema", [{}, {"type": "object", "properties": {"query": {"type": "string"}}}])
def test_translate_anthropic_client_web_search_preserves_schema_and_choice(schema: dict[str, object]) -> None:
from litellm.types.llms.anthropic import AnthropicMessagesRequest
request: Final = AnthropicMessagesRequest(
model="gpt-5.4-mini",
max_tokens=128,
messages=[{"role": "user", "content": "Search for current news"}],
tools=[{"name": "web_search", "input_schema": schema}],
tool_choice={"type": "tool", "name": "web_search"},
)
translated, _ = LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(request)
assert "web_search_options" not in translated
assert translated["tools"] == [
{"type": "function", "function": {"name": "web_search", "parameters": schema}}
]
assert translated["tool_choice"] == {"type": "function", "function": {"name": "web_search"}}
def test_translate_anthropic_to_openai_with_web_search_tool():
"""
Test that Anthropic web search tools are converted to web_search_options parameter.

View file

@ -34,10 +34,12 @@ from pydantic import ValidationError
from litellm.proxy.db.db_url_settings import (
PG_SSL_REQUEST,
DatabaseURLSettings,
token_refresh_params_from_url,
translate_libpq_ssl_params,
unsupported_db_scheme,
unsupported_db_scheme_message,
)
from litellm.proxy.db.pgbouncer import PgBouncerPlan, PgBouncerSettings, plan_pgbouncer
from litellm.proxy.db.token_auth import AzureEntraTokenAuth, RdsIamTokenAuth
@ -51,6 +53,8 @@ _MANAGED_DB_ENV_VARS = (
"AZURE_POSTGRESQL_AUTH",
"DATABASE_DISABLE_PREPARED_STATEMENTS",
"DATABASE_MAX_IDLE_CONNECTION_LIFETIME",
"DATABASE_SSLMODE",
"DATABASE_SSLROOTCERT",
"DATABASE_URL",
"DIRECT_URL",
"DATABASE_URL_READ_REPLICA",
@ -781,6 +785,109 @@ def test_libpq_verify_full_and_sslrootcert_become_prisma_strict_sslcert(monkeypa
}
def _tls_env(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("DATABASE_HOST", "writer.example.com")
monkeypatch.setenv("DATABASE_USER", "litellm")
monkeypatch.setenv("DATABASE_NAME", "litellm_db")
monkeypatch.setenv("DATABASE_SSLMODE", "verify-full")
monkeypatch.setenv("DATABASE_SSLROOTCERT", "/certs/rds-bundle.pem")
def test_tls_env_vars_make_the_minted_iam_writer_url_verify_the_server(monkeypatch: pytest.MonkeyPatch):
"""The supervisor starts PgBouncer from the URL assembled here, before any
config.yaml is read, so an IAM URL with no TLS params leaves PgBouncer on
``prefer`` (no SNI, no verification) and the RDS handshake fails."""
monkeypatch.setenv("IAM_TOKEN_DB_AUTH", "true")
_tls_env(monkeypatch)
with _stub_iam_token("WRITER_TOKEN"):
assert _apply() is True
url: Final = os.environ["DATABASE_URL"]
assert url.startswith("postgresql://litellm:WRITER_TOKEN@writer.example.com:5432/litellm_db?")
assert _query(url) == {
"sslmode": ["require"],
"sslcert": ["/certs/rds-bundle.pem"],
"sslaccept": ["strict"],
"max_idle_connection_lifetime": ["60"],
}
def test_tls_env_vars_apply_to_the_password_writer_and_the_assembled_reader(monkeypatch: pytest.MonkeyPatch):
_tls_env(monkeypatch)
monkeypatch.setenv("DATABASE_PASSWORD", "s3cr3t")
monkeypatch.setenv("DATABASE_SCHEMA", "public")
monkeypatch.setenv("DATABASE_HOST_READ_REPLICA", "reader.example.com")
assert _apply() is True
expected: Final = {
"schema": ["public"],
"sslmode": ["require"],
"sslcert": ["/certs/rds-bundle.pem"],
"sslaccept": ["strict"],
"max_idle_connection_lifetime": ["60"],
}
assert os.environ["DATABASE_URL"].startswith("postgresql://litellm:s3cr3t@writer.example.com:5432/litellm_db?")
assert _query(os.environ["DATABASE_URL"]) == expected
assert os.environ["DATABASE_URL_READ_REPLICA"].startswith(
"postgresql://litellm:s3cr3t@reader.example.com:5432/litellm_db?"
)
assert _query(os.environ["DATABASE_URL_READ_REPLICA"]) == expected
def test_sslrootcert_env_var_alone_means_verify_full_for_prisma_and_pgbouncer(monkeypatch: pytest.MonkeyPatch):
"""Under libpq's default ``prefer`` a root cert is never consulted, so a URL
carrying only ``sslrootcert`` would leave PgBouncer on ``prefer`` with the CA
loaded but unused. Supplying a CA and nothing else must verify."""
_tls_env(monkeypatch)
monkeypatch.delenv("DATABASE_SSLMODE")
monkeypatch.setenv("DATABASE_PASSWORD", "s3cr3t")
assert _apply() is True
url: Final = os.environ["DATABASE_URL"]
assert _query(url) == {
"sslmode": ["require"],
"sslcert": ["/certs/rds-bundle.pem"],
"sslaccept": ["strict"],
"max_idle_connection_lifetime": ["60"],
}
plan: Final = plan_pgbouncer(url, PgBouncerSettings(enabled=True), Path("/run/pgb"), None)
assert isinstance(plan, PgBouncerPlan), plan
assert "server_tls_sslmode = verify-full" in plan.ini
assert "server_tls_ca_file = /run/pgb/server-ca.pem" in plan.ini
def test_tls_env_vars_never_override_a_pinned_database_url(monkeypatch: pytest.MonkeyPatch):
writer: Final = (
"postgresql://pinned:url@db.example.com:5432/litellm_db?sslmode=disable&max_idle_connection_lifetime=60"
)
reader: Final = "postgresql://pinned:url@reader.example.com:5432/litellm_db?max_idle_connection_lifetime=60"
monkeypatch.setenv("DATABASE_URL", writer)
monkeypatch.setenv("DATABASE_URL_READ_REPLICA", reader)
monkeypatch.setenv("DATABASE_HOST_READ_REPLICA", "reader.example.com")
_tls_env(monkeypatch)
assert _apply() is False
assert os.environ["DATABASE_URL"] == writer
assert os.environ["DATABASE_URL_READ_REPLICA"] == reader
def test_token_refresh_params_keep_the_prisma_tls_dialect_but_not_the_schema():
kept: Final = token_refresh_params_from_url(
"postgresql://u:TOKEN@db.example.com:5432/litellm_db"
"?schema=tenant&connection_limit=5&sslmode=require&sslcert=/certs/root.pem&sslaccept=strict"
)
assert dict(kept) == {
"connection_limit": "5",
"sslmode": "require",
"sslcert": "/certs/root.pem",
"sslaccept": "strict",
}
def _issue_cert(
subject: str, issuer: x509.Certificate | None, issuer_key: ec.EllipticCurvePrivateKey | None, ca: bool
) -> tuple[x509.Certificate, ec.EllipticCurvePrivateKey]:

View file

@ -299,6 +299,10 @@ def test_azure_entra_mint_writes_an_encoded_url_into_the_db_url_env_var(azure_en
"connection_limit=20&pgbouncer=true&max_idle_connection_lifetime=45",
{"connection_limit": ["20"], "pgbouncer": ["true"], "max_idle_connection_lifetime": ["45"]},
),
(
"sslmode=require&sslcert=/certs/root.pem&sslaccept=strict&schema=tenant",
{"sslmode": ["require"], "sslcert": ["/certs/root.pem"], "sslaccept": ["strict"]},
),
],
)
def test_token_refresh_keeps_the_connection_params_of_the_url_it_replaces(

View file

@ -2,8 +2,9 @@ import asyncio
import logging
import os
import sys
from typing import Any, Dict
from unittest.mock import AsyncMock, MagicMock, patch
from typing import Any, Dict, Final
from unittest.mock import AsyncMock, MagicMock, call, patch
from urllib.parse import parse_qs, urlsplit
import pytest
@ -927,6 +928,47 @@ def test_prisma_client_init_falls_back_to_writer_when_reader_iam_token_fails(
)
def test_prisma_client_init_keeps_reader_tls_params_on_the_minted_iam_url(
monkeypatch: pytest.MonkeyPatch,
):
"""The initial reader mint rebuilds the URL from host/port/user/db, so the
Prisma TLS dialect on DATABASE_URL_READ_REPLICA must be carried over or
a verify-only database rejects the reader and reads fall to the writer."""
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
monkeypatch.setenv("IAM_TOKEN_DB_AUTH", "true")
monkeypatch.setenv(
"DATABASE_URL_READ_REPLICA",
"postgresql://reader_user@reader.aurora.local:5432/litellm"
"?schema=tenant&sslmode=require&sslcert=/certs/root.pem&sslaccept=strict",
)
prisma_factory: Final = MagicMock(name="Prisma")
fake_prisma_module: Final = MagicMock(Prisma=prisma_factory)
monkeypatch.setitem(sys.modules, "prisma", fake_prisma_module)
fake_iam_module: Final = MagicMock(generate_iam_auth_token=MagicMock(return_value="READER-TOKEN"))
monkeypatch.setitem(sys.modules, "litellm.proxy.auth.rds_iam_token", fake_iam_module)
from litellm.proxy.utils import PrismaClient
client: Final = PrismaClient(
database_url="postgresql://writer@writer.aurora.local:5432/litellm",
proxy_logging_obj=MagicMock(),
)
assert isinstance(client.db, RoutingPrismaWrapper)
reader_url: Final = os.environ["DATABASE_URL_READ_REPLICA"]
assert reader_url.startswith("postgresql://reader_user:READER-TOKEN@reader.aurora.local:5432/litellm?")
assert parse_qs(urlsplit(reader_url).query) == {
"schema": ["tenant"],
"sslmode": ["require"],
"sslcert": ["/certs/root.pem"],
"sslaccept": ["strict"],
}
assert prisma_factory.call_args_list == [call(), call(datasource={"url": reader_url})]
@pytest.mark.asyncio
async def test_connect_degrades_writer_when_reader_available():
"""A writer connect failure with a healthy reader must NOT abort proxy

View file

@ -187,6 +187,22 @@ def test_construct_database_url_from_env_vars_with_schema(monkeypatch):
}
def test_construct_database_url_from_env_vars_carries_tls_env(monkeypatch: pytest.MonkeyPatch):
"""The CLI password path builds its URL here, so DATABASE_SSLMODE and
DATABASE_SSLROOTCERT must reach PgBouncer through it too."""
monkeypatch.setenv("DATABASE_HOST", "db.example.com")
monkeypatch.setenv("DATABASE_USERNAME", "user")
monkeypatch.setenv("DATABASE_PASSWORD", "pass")
monkeypatch.setenv("DATABASE_NAME", "litellm")
monkeypatch.setenv("DATABASE_SCHEMA", "public")
monkeypatch.setenv("DATABASE_SSLMODE", "verify-full")
monkeypatch.setenv("DATABASE_SSLROOTCERT", "/etc/ssl/certs/ca-certificates.crt")
assert construct_database_url_from_env_vars() == (
"postgresql://user:pass@db.example.com/litellm"
"?schema=public&sslmode=verify-full&sslrootcert=%2Fetc%2Fssl%2Fcerts%2Fca-certificates.crt"
)
def test_construct_database_url_from_env_vars_error_path_missing_host(monkeypatch):
monkeypatch.delenv("DATABASE_HOST", raising=False)
monkeypatch.setenv("DATABASE_USERNAME", "user")