mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
fix(proxy): restore pre-config-wins handling of pass-through endpoints (#43962)
* fix(proxy): restore pre-config-wins handling of pass-through endpoints Config-wins (#41779) made general_settings.pass_through_endpoints a config-owned key. The DB reader then got the config list back as if it were DB rows, re-registered each entry without forward_headers on every DB sync, and the stripped copy won the route lookup, so a config pass-through with forward_headers: true stopped forwarding Authorization. UI create, update and delete of pass-throughs were also rejected while the config declared any. This puts pass-throughs back on their pre-#41779 path: the settings store no longer lets the config own the key, the config list is captured env-resolved at load_config, each DB sync merges DB entries with config entries on paths the DB does not declare, and /config/field/info reads the stored rows only. A UI pass-through write re-applies that merge immediately so the config entries stay served until the next sync. * fix(proxy): keep config pass-throughs in every reload of the merged list get_config now returns DB pass-throughs plus config ones on other paths, each DB sync republishes that merged list, and /config/field/info reads pass_through_endpoints from the DB row so a UI write never drops stored entries when models are not stored in the DB * fix(proxy): keep serving pass-throughs while the config file reloads load_yaml cleared the runtime pass-through list, so auth: false routes answered 401 while get_config awaited the database * fix(proxy): read stored pass-throughs from the writer before a UI write A lagging read replica could return an older list, and the UI create and edit flows write the whole field back * fix(proxy): apply config file pass-through auth changes on reload The kept runtime list was merged as if it were DB entries, so an edited config entry on the same path was dropped. Merge the stored DB row with the fresh config instead, and give the field-info test mock a writer * fix(proxy): keep pass-throughs served while a DB sync reads the database get_config resets the stored DB rows before reading them again, which cleared the served pass-through list and made auth: false routes answer 401 for the length of the read * refactor(proxy): move the settings store reload out of the loop basedpyright rejects a Final variable assigned inside a loop
This commit is contained in:
parent
2b19ddb7a3
commit
2eb2bf130b
8 changed files with 671 additions and 181 deletions
|
|
@ -78,6 +78,13 @@ def _build_dual_source_keys() -> Mapping[tuple[Section, str], KeyRule]:
|
|||
DUAL_SOURCE_KEYS: Final[Mapping[tuple[Section, str], KeyRule]] = _build_dual_source_keys()
|
||||
|
||||
|
||||
RESOURCE_LIST_KEYS: Final[frozenset[tuple[Section, str]]] = frozenset({("general_settings", "pass_through_endpoints")})
|
||||
|
||||
|
||||
def is_resource_list(section: Section, key: str) -> bool:
|
||||
return (section, key) in RESOURCE_LIST_KEYS
|
||||
|
||||
|
||||
def rule_for(section: Section, key: str) -> KeyRule:
|
||||
return DUAL_SOURCE_KEYS.get((section, key), DUAL_SOURCE_KEYS[(section, "*")])
|
||||
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from litellm.proxy.config_resolvers.settings_rules import (
|
|||
Resolved,
|
||||
Section,
|
||||
SettingValue,
|
||||
is_resource_list,
|
||||
resolve,
|
||||
rule_for,
|
||||
)
|
||||
|
|
@ -49,8 +50,13 @@ class SettingsStore(MutableMapping[str, JsonValue]):
|
|||
self._deleted_runtime_keys: frozenset[str] = frozenset()
|
||||
|
||||
def load_yaml(self, mapping: Mapping[str, JsonValue]) -> None:
|
||||
self._yaml_values = MappingProxyType(dict(mapping))
|
||||
self._clear_runtime()
|
||||
self._yaml_values = MappingProxyType(
|
||||
{key: value for key, value in mapping.items() if not is_resource_list(self._section, key)}
|
||||
)
|
||||
self._runtime_values = MappingProxyType(
|
||||
{key: value for key, value in self._runtime_values.items() if is_resource_list(self._section, key)}
|
||||
)
|
||||
self._deleted_runtime_keys = frozenset()
|
||||
|
||||
def config_value(self, key: str) -> JsonValue:
|
||||
return self._yaml_values.get(key)
|
||||
|
|
@ -136,10 +142,6 @@ class SettingsStore(MutableMapping[str, JsonValue]):
|
|||
def __bool__(self) -> bool:
|
||||
return any(True for _ in self)
|
||||
|
||||
def _clear_runtime(self) -> None:
|
||||
self._runtime_values = _EMPTY_VALUES
|
||||
self._deleted_runtime_keys = frozenset()
|
||||
|
||||
def _clear_runtime_keys(self, keys: frozenset[str]) -> None:
|
||||
stale: Final = frozenset(key for key in keys if not self.owned_by_config(key))
|
||||
if not stale:
|
||||
|
|
@ -160,6 +162,9 @@ class SettingsStore(MutableMapping[str, JsonValue]):
|
|||
)
|
||||
)
|
||||
|
||||
def db_value(self, key: str) -> SettingValue:
|
||||
return self._db_value(key) if is_resource_list(self._section, key) else ABSENT
|
||||
|
||||
def _db_value(self, key: str) -> SettingValue:
|
||||
rule: Final = rule_for(self._section, key)
|
||||
return self._database_rows.get(rule.db_row, _EMPTY_VALUES).get(key, ABSENT)
|
||||
|
|
|
|||
|
|
@ -500,9 +500,13 @@ from litellm.proxy.config_resolvers.alerting import (
|
|||
)
|
||||
from litellm.proxy.config_resolvers.changed_section_keys import changed_section_keys
|
||||
from litellm.proxy.config_resolvers.settings_rules import (
|
||||
ABSENT,
|
||||
DbRow,
|
||||
Section,
|
||||
SettingValue,
|
||||
coerce_bool,
|
||||
is_absent,
|
||||
is_resource_list,
|
||||
)
|
||||
from litellm.proxy.config_resolvers.settings_rules import (
|
||||
JsonValue as SettingsJsonValue,
|
||||
|
|
@ -5218,6 +5222,8 @@ class _ConfigWithBaseline(dict[str, object]):
|
|||
|
||||
_EMPTY_SETTINGS_MAPPING: Final[Mapping[str, SettingsJsonValue]] = MappingProxyType({})
|
||||
_SETTINGS_MAPPING: Final = TypeAdapter(dict[str, SettingsJsonValue])
|
||||
_SETTINGS_LIST: Final = TypeAdapter(list[SettingsJsonValue])
|
||||
_ENDPOINT_DICTS: Final = TypeAdapter(list[dict[str, object]])
|
||||
|
||||
|
||||
def _as_settings_mapping(value: object) -> Mapping[str, SettingsJsonValue]:
|
||||
|
|
@ -5232,6 +5238,40 @@ def _get_field_default(field_info: FieldInfo) -> JsonValue:
|
|||
return cast(JsonValue, field_info.default) # cast-ok: Pydantic field defaults are JSON values at runtime
|
||||
|
||||
|
||||
def _pass_through_endpoints_beside_db(db_endpoints: object, config_endpoints: object) -> list[SettingsJsonValue]:
|
||||
stored: Final = db_endpoints if isinstance(db_endpoints, list) else ()
|
||||
declared: Final = config_endpoints if isinstance(config_endpoints, list) else ()
|
||||
db_paths: Final = frozenset(endpoint.get("path") for endpoint in stored if isinstance(endpoint, dict))
|
||||
beside_db: Final = (
|
||||
endpoint for endpoint in declared if not isinstance(endpoint, dict) or endpoint.get("path") not in db_paths
|
||||
)
|
||||
return _SETTINGS_LIST.validate_python((*stored, *beside_db))
|
||||
|
||||
|
||||
def _with_config_file_pass_through_endpoints(
|
||||
section_config: object, resolved: Mapping[str, SettingsJsonValue], db_endpoints: SettingValue
|
||||
) -> Mapping[str, object]:
|
||||
config_endpoints: Final = (
|
||||
section_config.get("pass_through_endpoints") if isinstance(section_config, Mapping) else None
|
||||
)
|
||||
if config_endpoints is None and not isinstance(db_endpoints, list) and "pass_through_endpoints" not in resolved:
|
||||
return resolved
|
||||
return MappingProxyType(
|
||||
{
|
||||
**resolved,
|
||||
"pass_through_endpoints": _pass_through_endpoints_beside_db(db_endpoints, config_endpoints),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _reload_settings_store(section: Section, store: SettingsStore, section_config: object) -> None:
|
||||
serving_pass_throughs: Final = store.get("pass_through_endpoints")
|
||||
store.load_yaml(_as_settings_mapping(section_config))
|
||||
store.apply_db_row(section, _EMPTY_SETTINGS_MAPPING)
|
||||
if is_resource_list(section, "pass_through_endpoints") and serving_pass_throughs is not None:
|
||||
store["pass_through_endpoints"] = serving_pass_throughs
|
||||
|
||||
|
||||
def _bind_general_settings_store(settings: SettingsStore) -> None:
|
||||
global general_settings
|
||||
general_settings = settings # pyright: ignore[reportAssignmentType] # legacy global accepts mappings
|
||||
|
|
@ -5364,22 +5404,18 @@ class ProxyConfig:
|
|||
)
|
||||
|
||||
def _load_yaml_settings_stores(self, config: Mapping[str, object]) -> None:
|
||||
global config_passthrough_endpoints
|
||||
for section, store in self._settings_stores.items():
|
||||
store.load_yaml(_as_settings_mapping(config.get(section)))
|
||||
store.apply_db_row(section, _EMPTY_SETTINGS_MAPPING)
|
||||
yaml_endpoints: Final = self.settings.config_value("pass_through_endpoints")
|
||||
config_passthrough_endpoints = (
|
||||
[dict(endpoint) for endpoint in yaml_endpoints if isinstance(endpoint, dict)]
|
||||
if isinstance(yaml_endpoints, list)
|
||||
else None
|
||||
)
|
||||
_reload_settings_store(section, store, config.get(section))
|
||||
|
||||
def _config_with_resolved_settings(self, config: Mapping[str, object]) -> dict[str, object]:
|
||||
return { # mutable-ok: get_config preserves the mutable mapping contract used by existing loaders
|
||||
**config,
|
||||
**{
|
||||
section: dict(store.resolved())
|
||||
section: dict(
|
||||
_with_config_file_pass_through_endpoints(
|
||||
config.get(section), store.resolved(), store.db_value("pass_through_endpoints")
|
||||
)
|
||||
)
|
||||
for section, store in self._settings_stores.items()
|
||||
if isinstance(config.get(section), Mapping) or len(store) > 0
|
||||
},
|
||||
|
|
@ -6743,6 +6779,7 @@ class ProxyConfig:
|
|||
|
||||
## pass through endpoints
|
||||
if general_settings.get("pass_through_endpoints", None) is not None:
|
||||
config_passthrough_endpoints = general_settings["pass_through_endpoints"]
|
||||
await initialize_pass_through_endpoints(
|
||||
pass_through_endpoints=general_settings["pass_through_endpoints"],
|
||||
config_file_path=config_file_path,
|
||||
|
|
@ -7758,14 +7795,12 @@ class ProxyConfig:
|
|||
self.settings.load_yaml(_as_settings_mapping(general_settings))
|
||||
cache_size_was_db: Final = self.settings.source("user_api_key_cache_max_size") == "db"
|
||||
previous_cleanup_schedule: Final = self._resolved_cleanup_schedule()
|
||||
previous_pass_through_endpoints: Final = self.settings.get("pass_through_endpoints")
|
||||
self.settings.apply_db_row("general_settings", db_general_settings)
|
||||
_bind_general_settings_store(self.settings)
|
||||
await self._apply_general_settings_side_effects(
|
||||
db_general_settings,
|
||||
cache_size_was_db,
|
||||
previous_cleanup_schedule,
|
||||
previous_pass_through_endpoints,
|
||||
)
|
||||
|
||||
def _resolved_cleanup_schedule(self) -> tuple[object, ...]:
|
||||
|
|
@ -7779,11 +7814,10 @@ class ProxyConfig:
|
|||
db_values: Mapping[str, SettingsJsonValue],
|
||||
cache_size_was_db: bool,
|
||||
previous_cleanup_schedule: tuple[object, ...],
|
||||
previous_pass_through_endpoints: SettingsJsonValue | None,
|
||||
) -> None:
|
||||
effects: Final = (
|
||||
self._apply_alerting_settings,
|
||||
partial(self._apply_pass_through_settings, previous_endpoints=previous_pass_through_endpoints),
|
||||
self._apply_pass_through_settings,
|
||||
self._apply_boolean_settings,
|
||||
partial(self._apply_cache_size_setting, cache_size_was_db=cache_size_was_db),
|
||||
self._apply_store_model_in_db_setting,
|
||||
|
|
@ -7816,19 +7850,23 @@ class ProxyConfig:
|
|||
if "plugins" in db_values and self.settings.source("plugins") == "db":
|
||||
register_plugins_from_config(self.settings)
|
||||
|
||||
async def _apply_pass_through_settings(
|
||||
self,
|
||||
db_values: Mapping[str, SettingsJsonValue],
|
||||
previous_endpoints: SettingsJsonValue | None,
|
||||
) -> None:
|
||||
del db_values
|
||||
resolved_endpoints: Final = self.settings.get("pass_through_endpoints")
|
||||
if resolved_endpoints == previous_endpoints:
|
||||
async def _apply_pass_through_settings(self, db_values: Mapping[str, SettingsJsonValue]) -> None:
|
||||
db_endpoints: Final = db_values.get("pass_through_endpoints")
|
||||
if isinstance(db_endpoints, list):
|
||||
await self._serve_pass_through_endpoints(db_endpoints)
|
||||
return
|
||||
await initialize_pass_through_endpoints(
|
||||
pass_through_endpoints=resolved_endpoints if isinstance(resolved_endpoints, list) else []
|
||||
if "pass_through_endpoints" not in self.settings:
|
||||
self._publish_pass_through_endpoints(())
|
||||
|
||||
def _publish_pass_through_endpoints(self, db_endpoints: Sequence[SettingsJsonValue]) -> None:
|
||||
self.settings["pass_through_endpoints"] = _pass_through_endpoints_beside_db(
|
||||
list(db_endpoints), config_passthrough_endpoints
|
||||
)
|
||||
|
||||
async def _serve_pass_through_endpoints(self, db_endpoints: Sequence[SettingsJsonValue]) -> None:
|
||||
self._publish_pass_through_endpoints(db_endpoints)
|
||||
await initialize_pass_through_endpoints(pass_through_endpoints=_ENDPOINT_DICTS.validate_python(db_endpoints))
|
||||
|
||||
async def _apply_boolean_settings(self, db_values: Mapping[str, SettingsJsonValue]) -> None:
|
||||
for key in (
|
||||
"store_prompts_in_spend_logs",
|
||||
|
|
@ -18312,6 +18350,9 @@ async def update_config_general_settings(
|
|||
)
|
||||
await invalidate_config_param("general_settings")
|
||||
proxy_config.settings.apply_db_row("general_settings", general_settings)
|
||||
if is_resource_list("general_settings", data.field_name):
|
||||
stored_endpoints: Final = general_settings.get("pass_through_endpoints")
|
||||
await proxy_config._serve_pass_through_endpoints(stored_endpoints if isinstance(stored_endpoints, list) else ())
|
||||
asyncio.create_task(
|
||||
create_config_audit_log(
|
||||
"general_settings", "updated", before_general_settings, general_settings, user_api_key_dict
|
||||
|
|
@ -18463,6 +18504,20 @@ def _apply_webhook_role_gate(webhook_map, is_full_admin: bool):
|
|||
return {alert_type: "REDACTED" for alert_type in webhook_map}
|
||||
|
||||
|
||||
async def _declared_general_setting(
|
||||
settings: SettingsStore, field_name: str, prisma_client: PrismaClient
|
||||
) -> SettingValue:
|
||||
if is_resource_list("general_settings", field_name):
|
||||
row: Final = await ConfigRepository(prisma_client, use_writer=True).table.find_first(
|
||||
where={"param_name": "general_settings"}
|
||||
)
|
||||
stored: Final = row.param_value if row is not None and isinstance(row.param_value, Mapping) else {}
|
||||
return stored.get(field_name, ABSENT) if stored.get(field_name) is not None else ABSENT
|
||||
if field_name not in settings:
|
||||
return ABSENT
|
||||
return settings.config_value(field_name) if settings.owned_by_config(field_name) else settings[field_name]
|
||||
|
||||
|
||||
@router.get(
|
||||
"/config/field/info",
|
||||
tags=["config.yaml"],
|
||||
|
|
@ -18501,15 +18556,12 @@ async def get_config_general_settings(
|
|||
)
|
||||
|
||||
settings: Final = proxy_config.settings
|
||||
if field_name not in settings:
|
||||
declared: Final = await _declared_general_setting(settings, field_name, prisma_client)
|
||||
if is_absent(declared):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={"error": f"Field name={field_name} is not set"},
|
||||
)
|
||||
|
||||
declared: Final = (
|
||||
settings.config_value(field_name) if settings.owned_by_config(field_name) else settings[field_name]
|
||||
)
|
||||
field_value = _redact_general_setting_value(
|
||||
field_name,
|
||||
declared,
|
||||
|
|
@ -18920,6 +18972,9 @@ async def delete_config_general_settings(
|
|||
)
|
||||
await invalidate_config_param("general_settings")
|
||||
proxy_config.settings.apply_db_row("general_settings", general_settings)
|
||||
if is_resource_list("general_settings", data.field_name):
|
||||
stored_endpoints: Final = general_settings.get("pass_through_endpoints")
|
||||
await proxy_config._serve_pass_through_endpoints(stored_endpoints if isinstance(stored_endpoints, list) else ())
|
||||
asyncio.create_task(
|
||||
create_config_audit_log(
|
||||
"general_settings", "deleted", before_general_settings, general_settings, user_api_key_dict
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from litellm.proxy.config_resolvers.settings_rules import (
|
|||
Section,
|
||||
SettingValue,
|
||||
is_absent,
|
||||
is_resource_list,
|
||||
resolve,
|
||||
rule_for,
|
||||
)
|
||||
|
|
@ -88,7 +89,6 @@ _PREVIOUSLY_DB_WINS: Final[tuple[str, ...]] = (
|
|||
"user_url_allowed_hosts",
|
||||
"provider_url_destination_allowed_hosts",
|
||||
"alerting",
|
||||
"pass_through_endpoints",
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -105,8 +105,9 @@ def test_the_store_resolves_every_config_and_stored_value_combination(
|
|||
section: Section, key: str, config_value: SettingValue, db_value: SettingValue
|
||||
) -> None:
|
||||
store: Final = _store_for(section, key, config_value, db_value)
|
||||
owned_config_value: Final = ABSENT if is_resource_list(section, key) else config_value
|
||||
|
||||
if not is_absent(config_value):
|
||||
if not is_absent(owned_config_value):
|
||||
assert store[key] == config_value
|
||||
assert store.source(key) == "config"
|
||||
elif is_absent(db_value) or db_value is None:
|
||||
|
|
@ -121,7 +122,7 @@ def test_the_store_resolves_every_config_and_stored_value_combination(
|
|||
def test_the_store_and_the_resolver_never_disagree(
|
||||
section: Section, key: str, config_value: SettingValue, db_value: SettingValue
|
||||
) -> None:
|
||||
resolved: Final = resolve(config_value, db_value)
|
||||
resolved: Final = resolve(ABSENT if is_resource_list(section, key) else config_value, db_value)
|
||||
store: Final = _store_for(section, key, config_value, db_value)
|
||||
|
||||
assert store.source(key) == resolved.source
|
||||
|
|
|
|||
|
|
@ -302,6 +302,28 @@ async def test_load_config_returns_and_binds_the_general_settings_store(tmp_path
|
|||
assert config_state["general_settings"]["max_file_size_mb"] == 5
|
||||
|
||||
|
||||
def test_settings_store_leaves_pass_through_endpoints_to_the_database() -> None:
|
||||
store: Final = SettingsStore("general_settings")
|
||||
store.load_yaml({"pass_through_endpoints": [{"path": "/config"}]})
|
||||
store.apply_db_row("general_settings", {"pass_through_endpoints": [{"path": "/db"}]})
|
||||
|
||||
assert store["pass_through_endpoints"] == [{"path": "/db"}]
|
||||
assert store.source("pass_through_endpoints") == "db"
|
||||
assert store.rejected_writes({"pass_through_endpoints": [{"path": "/ui"}]}) == ()
|
||||
|
||||
|
||||
def test_settings_store_keeps_serving_pass_through_endpoints_while_the_config_file_reloads() -> None:
|
||||
store: Final = SettingsStore("general_settings")
|
||||
store.load_yaml({"pass_through_endpoints": [{"path": "/config"}], "max_parallel_requests": 1})
|
||||
store["pass_through_endpoints"] = [{"path": "/config", "auth": False}]
|
||||
store["allowed_ips"] = ["1.2.3.4"]
|
||||
|
||||
store.load_yaml({"pass_through_endpoints": [{"path": "/config"}], "max_parallel_requests": 1})
|
||||
|
||||
assert store["pass_through_endpoints"] == [{"path": "/config", "auth": False}]
|
||||
assert "allowed_ips" not in store
|
||||
|
||||
|
||||
def test_settings_store_starts_with_an_unset_source() -> None:
|
||||
store: Final = SettingsStore("general_settings")
|
||||
|
||||
|
|
|
|||
|
|
@ -7683,3 +7683,534 @@ def test_passthrough_attributes_a_cli_session_to_its_alias_not_the_login_token()
|
|||
metadata = kwargs["litellm_params"]["metadata"]
|
||||
assert metadata["user_api_key"] == "cli-session-alice"
|
||||
assert _get_spend_logs_metadata(metadata)["user_api_key"] == "cli-session-alice"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _StoredConfigRow:
|
||||
param_name: str
|
||||
param_value: Mapping[str, object]
|
||||
|
||||
|
||||
class _InMemoryConfigTable:
|
||||
def __init__(self, rows: Mapping[str, Mapping[str, object]]) -> None:
|
||||
self.rows: dict[str, Mapping[str, object]] = dict(rows)
|
||||
self.db: Final = SimpleNamespace(litellm_config=self)
|
||||
self.writer_db: Final = SimpleNamespace(litellm_config=self)
|
||||
|
||||
def _row(self, param_name: str) -> _StoredConfigRow | None:
|
||||
value: Final = self.rows.get(param_name)
|
||||
return None if value is None else _StoredConfigRow(param_name=param_name, param_value=value)
|
||||
|
||||
async def get_generic_data(self, key: str, value: str, table_name: str) -> _StoredConfigRow | None:
|
||||
return self._row(value)
|
||||
|
||||
async def find_first(self, where: Mapping[str, str]) -> _StoredConfigRow | None:
|
||||
return self._row(where["param_name"])
|
||||
|
||||
async def find_unique(self, where: Mapping[str, str]) -> _StoredConfigRow | None:
|
||||
return self._row(where["param_name"])
|
||||
|
||||
async def upsert(self, where: Mapping[str, str], data: Mapping[str, Mapping[str, str]]) -> _StoredConfigRow:
|
||||
self.rows[where["param_name"]] = json.loads(data["update"]["param_value"])
|
||||
return _StoredConfigRow(param_name=where["param_name"], param_value=self.rows[where["param_name"]])
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _DbBackedProxy:
|
||||
proxy_config: object
|
||||
config_path: str
|
||||
config_table: _InMemoryConfigTable
|
||||
|
||||
|
||||
async def _boot_db_backed_proxy(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
config_pass_through_endpoints: list[dict[str, object]],
|
||||
db_pass_through_endpoints: list[dict[str, object]],
|
||||
master_key: str | None = None,
|
||||
store_model_in_db: bool = True,
|
||||
) -> _DbBackedProxy:
|
||||
import yaml
|
||||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy import proxy_server
|
||||
from litellm.proxy import utils as proxy_utils
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import _registered_pass_through_routes
|
||||
|
||||
general_settings: Final[dict[str, object]] = {"pass_through_endpoints": config_pass_through_endpoints}
|
||||
if master_key is not None:
|
||||
general_settings["master_key"] = master_key
|
||||
config_path: Final = tmp_path / "config.yaml"
|
||||
config_path.write_text(yaml.safe_dump({"model_list": [], "general_settings": general_settings}))
|
||||
config_table: Final = _InMemoryConfigTable(
|
||||
{"general_settings": {"pass_through_endpoints": db_pass_through_endpoints}} if db_pass_through_endpoints else {}
|
||||
)
|
||||
proxy_config: Final = proxy_server.ProxyConfig()
|
||||
monkeypatch.setattr(proxy_server, "proxy_config", proxy_config)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
monkeypatch.setattr(proxy_server, "user_config_file_path", str(config_path))
|
||||
monkeypatch.setattr(proxy_server, "general_settings", {})
|
||||
monkeypatch.setattr(proxy_server, "config_passthrough_endpoints", None)
|
||||
monkeypatch.setattr(proxy_server, "master_key", None)
|
||||
monkeypatch.setattr(proxy_server, "premium_user", False)
|
||||
monkeypatch.setattr(proxy_utils, "litellm_config_cache", DualCache())
|
||||
monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False)
|
||||
monkeypatch.delitem(proxy_server.app.dependency_overrides, user_api_key_auth, raising=False)
|
||||
_registered_pass_through_routes.clear()
|
||||
|
||||
await proxy_config.load_config(router=None, config_file_path=str(config_path))
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", config_table)
|
||||
monkeypatch.setattr(proxy_server, "store_model_in_db", store_model_in_db)
|
||||
return _DbBackedProxy(proxy_config, str(config_path), config_table)
|
||||
|
||||
|
||||
async def _run_db_sync_cycle(proxy: _DbBackedProxy) -> None:
|
||||
await proxy.proxy_config.get_config(config_file_path=proxy.config_path)
|
||||
await proxy.proxy_config._update_general_settings(proxy.config_table.rows.get("general_settings", {}))
|
||||
await proxy.proxy_config._init_pass_through_endpoints_in_db()
|
||||
|
||||
|
||||
async def _send_through_proxy(
|
||||
path: str, headers: Mapping[str, str], method: str = "POST"
|
||||
) -> tuple[httpx.Response, list[httpx.Request]]:
|
||||
from litellm.proxy.proxy_server import app
|
||||
|
||||
upstream_requests: Final[list[httpx.Request]] = []
|
||||
|
||||
def upstream(request: httpx.Request) -> httpx.Response:
|
||||
upstream_requests.append(request)
|
||||
return httpx.Response(200, json={"ok": True}, request=request)
|
||||
|
||||
fake_client, cleanup = _inject_fake_passthrough_client(httpx.MockTransport(upstream), timeout=None)
|
||||
try:
|
||||
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://proxy.test") as client:
|
||||
response = await client.request(method, path, headers=dict(headers), json={"q": 1})
|
||||
finally:
|
||||
cleanup()
|
||||
await fake_client.aclose()
|
||||
return response, upstream_requests
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_config_pass_through_keeps_forwarding_client_headers_after_a_db_sync(tmp_path, monkeypatch):
|
||||
proxy: Final = await _boot_db_backed_proxy(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
config_pass_through_endpoints=[
|
||||
{
|
||||
"path": "/cfg-forward",
|
||||
"target": "http://config-upstream.test/api",
|
||||
"forward_headers": True,
|
||||
"auth": False,
|
||||
}
|
||||
],
|
||||
db_pass_through_endpoints=[],
|
||||
)
|
||||
await _run_db_sync_cycle(proxy)
|
||||
|
||||
response, upstream_requests = await _send_through_proxy("/cfg-forward", {"Authorization": "Bearer caller-jwt"})
|
||||
|
||||
assert response.status_code == 200
|
||||
assert [str(request.url) for request in upstream_requests] == ["http://config-upstream.test/api"]
|
||||
assert upstream_requests[0].headers["authorization"] == "Bearer caller-jwt"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_config_and_db_pass_throughs_both_serve_and_list_after_a_db_sync(tmp_path, monkeypatch):
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import get_pass_through_endpoints
|
||||
|
||||
proxy: Final = await _boot_db_backed_proxy(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
config_pass_through_endpoints=[
|
||||
{"path": "/cfg-only", "target": "http://config-upstream.test/api", "auth": False}
|
||||
],
|
||||
db_pass_through_endpoints=[
|
||||
{"id": "db-endpoint", "path": "/db-only", "target": "http://db-upstream.test/api", "auth": False}
|
||||
],
|
||||
)
|
||||
await _run_db_sync_cycle(proxy)
|
||||
|
||||
config_response, config_upstream = await _send_through_proxy("/cfg-only", {})
|
||||
db_response, db_upstream = await _send_through_proxy("/db-only", {})
|
||||
listed: Final = await get_pass_through_endpoints(
|
||||
endpoint_id=None,
|
||||
team_id=None,
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role="proxy_admin"),
|
||||
)
|
||||
|
||||
assert (config_response.status_code, db_response.status_code) == (200, 200)
|
||||
assert [str(request.url) for request in config_upstream] == ["http://config-upstream.test/api"]
|
||||
assert [str(request.url) for request in db_upstream] == ["http://db-upstream.test/api"]
|
||||
assert sorted((endpoint.path, endpoint.is_from_config) for endpoint in listed.endpoints) == [
|
||||
("/cfg-only", True),
|
||||
("/db-only", False),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"stored_after_delete",
|
||||
[{"pass_through_endpoints": []}, {}],
|
||||
ids=["emptied-list", "dropped-key"],
|
||||
)
|
||||
async def test_a_deleted_db_pass_through_stops_serving_on_the_next_db_sync(tmp_path, monkeypatch, stored_after_delete):
|
||||
proxy: Final = await _boot_db_backed_proxy(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
config_pass_through_endpoints=[
|
||||
{"path": "/cfg-kept", "target": "http://config-upstream.test/api", "auth": False}
|
||||
],
|
||||
db_pass_through_endpoints=[
|
||||
{"id": "db-gone", "path": "/db-gone", "target": "http://db-upstream.test/api", "auth": False}
|
||||
],
|
||||
master_key="sk-pass-through-master",
|
||||
)
|
||||
await _run_db_sync_cycle(proxy)
|
||||
served_before, _ = await _send_through_proxy("/db-gone", {})
|
||||
|
||||
proxy.config_table.rows["general_settings"] = stored_after_delete
|
||||
await _run_db_sync_cycle(proxy)
|
||||
served_after, db_upstream = await _send_through_proxy("/db-gone", {})
|
||||
config_after, config_upstream = await _send_through_proxy("/cfg-kept", {})
|
||||
|
||||
assert (served_before.status_code, served_after.status_code, config_after.status_code) == (200, 401, 200)
|
||||
assert db_upstream == []
|
||||
assert [str(request.url) for request in config_upstream] == ["http://config-upstream.test/api"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_config_pass_through_reads_its_custom_key_header_when_the_db_holds_pass_throughs(
|
||||
tmp_path, monkeypatch
|
||||
):
|
||||
proxy: Final = await _boot_db_backed_proxy(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
config_pass_through_endpoints=[
|
||||
{
|
||||
"path": "/cfg-keyed",
|
||||
"target": "http://config-upstream.test/api",
|
||||
"auth": True,
|
||||
"headers": {"litellm_user_api_key": "x-cfg-key"},
|
||||
}
|
||||
],
|
||||
db_pass_through_endpoints=[
|
||||
{"id": "db-endpoint", "path": "/db-only", "target": "http://db-upstream.test/api", "auth": False}
|
||||
],
|
||||
master_key="sk-pass-through-master",
|
||||
)
|
||||
await _run_db_sync_cycle(proxy)
|
||||
|
||||
response, upstream_requests = await _send_through_proxy("/cfg-keyed", {"x-cfg-key": "sk-pass-through-master"})
|
||||
|
||||
assert response.status_code == 200
|
||||
assert [str(request.url) for request in upstream_requests] == ["http://config-upstream.test/api"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ui_can_create_a_db_pass_through_when_the_config_declares_pass_throughs(tmp_path, monkeypatch):
|
||||
from litellm.proxy._types import PassThroughGenericEndpoint
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import create_pass_through_endpoints
|
||||
|
||||
proxy: Final = await _boot_db_backed_proxy(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
config_pass_through_endpoints=[
|
||||
{"path": "/cfg-only", "target": "http://config-upstream.test/api", "auth": False}
|
||||
],
|
||||
db_pass_through_endpoints=[],
|
||||
)
|
||||
await _run_db_sync_cycle(proxy)
|
||||
|
||||
await create_pass_through_endpoints(
|
||||
data=PassThroughGenericEndpoint(path="/ui-made", target="http://ui-upstream.test/api", auth=False),
|
||||
request=MagicMock(spec=Request),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role="proxy_admin"),
|
||||
)
|
||||
await _run_db_sync_cycle(proxy)
|
||||
response, upstream_requests = await _send_through_proxy("/ui-made", {})
|
||||
|
||||
assert [endpoint["path"] for endpoint in proxy.config_table.rows["general_settings"]["pass_through_endpoints"]] == [
|
||||
"/ui-made"
|
||||
]
|
||||
assert response.status_code == 200
|
||||
assert [str(request.url) for request in upstream_requests] == ["http://ui-upstream.test/api"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_ui_created_pass_through_leaves_the_config_ones_open_before_the_next_db_sync(tmp_path, monkeypatch):
|
||||
from litellm.proxy._types import PassThroughGenericEndpoint
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import create_pass_through_endpoints
|
||||
|
||||
proxy: Final = await _boot_db_backed_proxy(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
config_pass_through_endpoints=[
|
||||
{"path": "/cfg-open", "target": "http://config-upstream.test/api", "auth": False, "forward_headers": True}
|
||||
],
|
||||
db_pass_through_endpoints=[],
|
||||
master_key="sk-pass-through-master",
|
||||
)
|
||||
await _run_db_sync_cycle(proxy)
|
||||
|
||||
await create_pass_through_endpoints(
|
||||
data=PassThroughGenericEndpoint(path="/ui-open", target="http://ui-upstream.test/api", auth=False),
|
||||
request=MagicMock(spec=Request),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role="proxy_admin"),
|
||||
)
|
||||
config_response, config_upstream = await _send_through_proxy("/cfg-open", {"Authorization": "Bearer caller-jwt"})
|
||||
ui_response, ui_upstream = await _send_through_proxy("/ui-open", {})
|
||||
|
||||
assert (config_response.status_code, ui_response.status_code) == (200, 200)
|
||||
assert [request.headers.get("authorization") for request in config_upstream] == ["Bearer caller-jwt"]
|
||||
assert [str(request.url) for request in ui_upstream] == ["http://ui-upstream.test/api"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_config_pass_through_serves_right_after_boot(tmp_path, monkeypatch):
|
||||
await _boot_db_backed_proxy(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
config_pass_through_endpoints=[
|
||||
{"path": "/cfg-boot", "target": "http://config-upstream.test/api", "auth": False}
|
||||
],
|
||||
db_pass_through_endpoints=[],
|
||||
)
|
||||
|
||||
response, upstream_requests = await _send_through_proxy("/cfg-boot", {})
|
||||
|
||||
assert response.status_code == 200
|
||||
assert [str(request.url) for request in upstream_requests] == ["http://config-upstream.test/api"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_config_pass_through_resolves_an_os_environ_target(tmp_path, monkeypatch):
|
||||
monkeypatch.setenv("LIT_PASS_THROUGH_TEST_UPSTREAM", "http://env-upstream.test/api")
|
||||
proxy: Final = await _boot_db_backed_proxy(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
config_pass_through_endpoints=[
|
||||
{"path": "/cfg-env", "target": "os.environ/LIT_PASS_THROUGH_TEST_UPSTREAM", "auth": False}
|
||||
],
|
||||
db_pass_through_endpoints=[],
|
||||
)
|
||||
|
||||
at_boot, at_boot_upstream = await _send_through_proxy("/cfg-env", {})
|
||||
await _run_db_sync_cycle(proxy)
|
||||
after_sync, after_sync_upstream = await _send_through_proxy("/cfg-env", {})
|
||||
|
||||
assert (at_boot.status_code, after_sync.status_code) == (200, 200)
|
||||
assert [str(request.url) for request in (*at_boot_upstream, *after_sync_upstream)] == [
|
||||
"http://env-upstream.test/api",
|
||||
"http://env-upstream.test/api",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_settings_write_keeps_the_config_file_pass_throughs(tmp_path, monkeypatch):
|
||||
import yaml
|
||||
|
||||
proxy: Final = await _boot_db_backed_proxy(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
config_pass_through_endpoints=[
|
||||
{"path": "/cfg-kept", "target": "http://config-upstream.test/api", "auth": False}
|
||||
],
|
||||
db_pass_through_endpoints=[],
|
||||
store_model_in_db=False,
|
||||
)
|
||||
config: Final = await proxy.proxy_config.get_config(config_file_path=proxy.config_path)
|
||||
|
||||
await proxy.proxy_config.save_config(
|
||||
new_config={**config, "general_settings": {**config["general_settings"], "max_parallel_requests": 7}}
|
||||
)
|
||||
|
||||
saved_general_settings: Final = yaml.safe_load(open(proxy.config_path))["general_settings"]
|
||||
assert saved_general_settings["max_parallel_requests"] == 7
|
||||
assert [endpoint["path"] for endpoint in saved_general_settings["pass_through_endpoints"]] == ["/cfg-kept"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_config_reload_keeps_config_pass_throughs_open_next_to_db_ones(tmp_path, monkeypatch):
|
||||
proxy: Final = await _boot_db_backed_proxy(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
config_pass_through_endpoints=[
|
||||
{"path": "/cfg-open", "target": "http://config-upstream.test/api", "auth": False, "forward_headers": True}
|
||||
],
|
||||
db_pass_through_endpoints=[
|
||||
{"id": "db-endpoint", "path": "/db-only", "target": "http://db-upstream.test/api", "auth": False}
|
||||
],
|
||||
master_key="sk-pass-through-master",
|
||||
)
|
||||
await _run_db_sync_cycle(proxy)
|
||||
|
||||
await proxy.proxy_config.get_config(config_file_path=proxy.config_path)
|
||||
response, upstream_requests = await _send_through_proxy("/cfg-open", {"Authorization": "Bearer caller-jwt"})
|
||||
|
||||
assert response.status_code == 200
|
||||
assert [request.headers.get("authorization") for request in upstream_requests] == ["Bearer caller-jwt"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ui_create_keeps_the_stored_pass_throughs_when_models_are_not_stored_in_the_db(tmp_path, monkeypatch):
|
||||
from litellm.proxy._types import PassThroughGenericEndpoint
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import create_pass_through_endpoints
|
||||
|
||||
proxy: Final = await _boot_db_backed_proxy(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
config_pass_through_endpoints=[
|
||||
{"path": "/cfg-only", "target": "http://config-upstream.test/api", "auth": False}
|
||||
],
|
||||
db_pass_through_endpoints=[
|
||||
{"id": "db-endpoint", "path": "/db-stored", "target": "http://db-upstream.test/api", "auth": False}
|
||||
],
|
||||
store_model_in_db=False,
|
||||
)
|
||||
|
||||
await create_pass_through_endpoints(
|
||||
data=PassThroughGenericEndpoint(path="/ui-made", target="http://ui-upstream.test/api", auth=False),
|
||||
request=MagicMock(spec=Request),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role="proxy_admin"),
|
||||
)
|
||||
|
||||
assert [endpoint["path"] for endpoint in proxy.config_table.rows["general_settings"]["pass_through_endpoints"]] == [
|
||||
"/db-stored",
|
||||
"/ui-made",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_deleting_the_stored_pass_through_field_stops_serving_its_routes_right_away(tmp_path, monkeypatch):
|
||||
from litellm.proxy._types import ConfigFieldDelete
|
||||
from litellm.proxy.proxy_server import delete_config_general_settings
|
||||
|
||||
proxy: Final = await _boot_db_backed_proxy(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
config_pass_through_endpoints=[
|
||||
{"path": "/cfg-kept", "target": "http://config-upstream.test/api", "auth": False}
|
||||
],
|
||||
db_pass_through_endpoints=[
|
||||
{"id": "db-gone", "path": "/db-gone", "target": "http://db-upstream.test/api", "auth": False}
|
||||
],
|
||||
master_key="sk-pass-through-master",
|
||||
)
|
||||
await _run_db_sync_cycle(proxy)
|
||||
served_before, _ = await _send_through_proxy("/db-gone", {})
|
||||
|
||||
await delete_config_general_settings(
|
||||
data=ConfigFieldDelete(config_type="general_settings", field_name="pass_through_endpoints"),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role="proxy_admin"),
|
||||
)
|
||||
served_after, db_upstream = await _send_through_proxy("/db-gone", {})
|
||||
config_after, _ = await _send_through_proxy("/cfg-kept", {})
|
||||
|
||||
assert (served_before.status_code, served_after.status_code, config_after.status_code) == (200, 401, 200)
|
||||
assert db_upstream == []
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _LaggingReadReplica:
|
||||
writer: _InMemoryConfigTable
|
||||
|
||||
async def find_first(self, where: Mapping[str, str]) -> _StoredConfigRow | None:
|
||||
return None
|
||||
|
||||
async def upsert(self, where: Mapping[str, str], data: Mapping[str, Mapping[str, str]]) -> _StoredConfigRow:
|
||||
return await self.writer.upsert(where=where, data=data)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ui_create_keeps_stored_pass_throughs_a_lagging_read_replica_has_not_seen(tmp_path, monkeypatch):
|
||||
from litellm.proxy._types import PassThroughGenericEndpoint
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import create_pass_through_endpoints
|
||||
|
||||
proxy: Final = await _boot_db_backed_proxy(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
config_pass_through_endpoints=[],
|
||||
db_pass_through_endpoints=[
|
||||
{"id": "db-endpoint", "path": "/db-stored", "target": "http://db-upstream.test/api", "auth": False}
|
||||
],
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
proxy.config_table, "db", SimpleNamespace(litellm_config=_LaggingReadReplica(proxy.config_table))
|
||||
)
|
||||
|
||||
await create_pass_through_endpoints(
|
||||
data=PassThroughGenericEndpoint(path="/ui-made", target="http://ui-upstream.test/api", auth=False),
|
||||
request=MagicMock(spec=Request),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role="proxy_admin"),
|
||||
)
|
||||
|
||||
assert [endpoint["path"] for endpoint in proxy.config_table.rows["general_settings"]["pass_through_endpoints"]] == [
|
||||
"/db-stored",
|
||||
"/ui-made",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_config_reload_applies_auth_turned_on_for_a_config_pass_through(tmp_path, monkeypatch):
|
||||
import yaml
|
||||
|
||||
proxy: Final = await _boot_db_backed_proxy(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
config_pass_through_endpoints=[
|
||||
{"path": "/cfg-locked", "target": "http://config-upstream.test/api", "auth": False}
|
||||
],
|
||||
db_pass_through_endpoints=[],
|
||||
master_key="sk-pass-through-master",
|
||||
)
|
||||
await _run_db_sync_cycle(proxy)
|
||||
open_before, _ = await _send_through_proxy("/cfg-locked", {})
|
||||
|
||||
reloaded_config: Final = yaml.safe_load(open(proxy.config_path))
|
||||
reloaded_config["general_settings"]["pass_through_endpoints"][0]["auth"] = True
|
||||
open(proxy.config_path, "w").write(yaml.safe_dump(reloaded_config))
|
||||
await _run_db_sync_cycle(proxy)
|
||||
locked_after, upstream_requests = await _send_through_proxy("/cfg-locked", {})
|
||||
|
||||
assert (open_before.status_code, locked_after.status_code) == (200, 401)
|
||||
assert upstream_requests == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pass_throughs_stay_open_while_a_db_sync_reads_the_database(tmp_path, monkeypatch):
|
||||
proxy: Final = await _boot_db_backed_proxy(
|
||||
tmp_path,
|
||||
monkeypatch,
|
||||
config_pass_through_endpoints=[
|
||||
{"path": "/cfg-open", "target": "http://config-upstream.test/api", "auth": False}
|
||||
],
|
||||
db_pass_through_endpoints=[
|
||||
{"id": "db-endpoint", "path": "/db-open", "target": "http://db-upstream.test/api", "auth": False}
|
||||
],
|
||||
master_key="sk-pass-through-master",
|
||||
)
|
||||
await _run_db_sync_cycle(proxy)
|
||||
database_read_started: Final = asyncio.Event()
|
||||
release_database_read: Final = asyncio.Event()
|
||||
read_row: Final = proxy.config_table.get_generic_data
|
||||
|
||||
async def slow_read(key: str, value: str, table_name: str) -> _StoredConfigRow | None:
|
||||
database_read_started.set()
|
||||
await release_database_read.wait()
|
||||
return await read_row(key=key, value=value, table_name=table_name)
|
||||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy import utils as proxy_utils
|
||||
|
||||
monkeypatch.setattr(proxy_utils, "litellm_config_cache", DualCache())
|
||||
monkeypatch.setattr(proxy.config_table, "get_generic_data", slow_read)
|
||||
sync: Final = asyncio.create_task(proxy.proxy_config.get_config(config_file_path=proxy.config_path))
|
||||
await asyncio.wait_for(database_read_started.wait(), timeout=5)
|
||||
config_during_sync, _ = await _send_through_proxy("/cfg-open", {})
|
||||
db_during_sync, _ = await _send_through_proxy("/db-open", {})
|
||||
release_database_read.set()
|
||||
await sync
|
||||
|
||||
assert (config_during_sync.status_code, db_during_sync.status_code) == (200, 200)
|
||||
|
|
|
|||
|
|
@ -4418,15 +4418,13 @@ async def test_ProxyConfig__update_general_settings_dispatches_every_side_effect
|
|||
for name, handler in handlers:
|
||||
monkeypatch.setattr(pc, name, handler)
|
||||
|
||||
await pc._apply_general_settings_side_effects({}, False, (), None)
|
||||
await pc._apply_general_settings_side_effects({}, False, ())
|
||||
|
||||
for name, handler in handlers:
|
||||
if name == "_apply_cache_size_setting":
|
||||
handler.assert_awaited_once_with({}, cache_size_was_db=False)
|
||||
elif name == "_apply_retention_settings":
|
||||
handler.assert_awaited_once_with({}, previous_cleanup_schedule=())
|
||||
elif name == "_apply_pass_through_settings":
|
||||
handler.assert_awaited_once_with({}, previous_endpoints=None)
|
||||
else:
|
||||
handler.assert_awaited_once_with({})
|
||||
|
||||
|
|
@ -4492,7 +4490,7 @@ async def test_ProxyConfig__update_config_from_db_resolves_through_settings_stor
|
|||
"max_file_size_mb": 7,
|
||||
"max_parallel_requests": 3,
|
||||
"alerting": ["config"],
|
||||
"pass_through_endpoints": [{"path": "/config"}],
|
||||
"pass_through_endpoints": [{"path": "/db"}, {"path": "/config"}],
|
||||
"maximum_spend_logs_cleanup_batch_size": 10,
|
||||
}
|
||||
assert resolved["router_settings"] == {"fallbacks": ["config"], "num_retries": 1}
|
||||
|
|
@ -4523,19 +4521,6 @@ async def test_ProxyConfig__update_config_from_db_keeps_keys_the_config_file_omi
|
|||
assert pc.settings.source("max_parallel_requests") == "db"
|
||||
|
||||
|
||||
def test_ProxyConfig_load_yaml_settings_stores_keeps_db_endpoints_out_of_config_baseline():
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
pc = ProxyConfig()
|
||||
config_endpoint: Final = {"path": "/config", "target": "https://config.example"}
|
||||
db_endpoint: Final = {"id": "db-endpoint", "path": "/db", "target": "https://db.example"}
|
||||
|
||||
pc._load_yaml_settings_stores({"general_settings": {"pass_through_endpoints": [config_endpoint]}})
|
||||
pc.settings.apply_db_row("general_settings", {"pass_through_endpoints": [db_endpoint]})
|
||||
|
||||
assert proxy_server.config_passthrough_endpoints == [config_endpoint]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ProxyConfig_add_deployment_continues_after_null_pass_through_endpoints(monkeypatch):
|
||||
from litellm.proxy import proxy_server
|
||||
|
|
|
|||
|
|
@ -9,7 +9,6 @@ import socket
|
|||
import subprocess
|
||||
import time
|
||||
import types
|
||||
import uuid
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
|
@ -8111,13 +8110,10 @@ async def test_update_general_settings_keeps_yaml_pass_through_endpoints_next_to
|
|||
[(None, None), (["POST"], ["GET"])],
|
||||
ids=["all-methods", "disjoint-methods"],
|
||||
)
|
||||
async def test_update_general_settings_db_pass_through_endpoint_cannot_override_a_yaml_declared_path(
|
||||
async def test_update_general_settings_db_pass_through_endpoint_overrides_yaml_entry_on_the_same_path(
|
||||
db_methods: list[str] | None, yaml_methods: list[str] | None
|
||||
):
|
||||
"""``pass_through_endpoints`` is config-owned once the file declares it, so a stored
|
||||
``auth: true`` entry on a path the YAML already declares ``auth: false`` no longer
|
||||
locks that path down. Changing it means editing the config file. A path the YAML
|
||||
does not declare is still governed by the stored row, which the sibling test covers."""
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.proxy_server import ProxyConfig
|
||||
|
||||
yaml_endpoint: Final = {
|
||||
|
|
@ -8140,129 +8136,16 @@ async def test_update_general_settings_db_pass_through_endpoint_cannot_override_
|
|||
request.headers = {}
|
||||
request.query_params = {}
|
||||
|
||||
settings: Final = patch(
|
||||
"litellm.proxy.proxy_server.general_settings", {"pass_through_endpoints": [yaml_endpoint]}
|
||||
) # test-quality-ok: the method reads this module global; no injection seam
|
||||
yaml_endpoints: Final = patch(
|
||||
"litellm.proxy.proxy_server.config_passthrough_endpoints", [yaml_endpoint]
|
||||
) # test-quality-ok: module global holding the YAML endpoints the fix merges in
|
||||
initialize: Final = patch(
|
||||
"litellm.proxy.proxy_server.initialize_pass_through_endpoints", AsyncMock()
|
||||
) # test-quality-ok: route registration needs the FastAPI app; auth is the observable here
|
||||
master_key: Final = patch(
|
||||
"litellm.proxy.proxy_server.master_key", "sk-master"
|
||||
) # test-quality-ok: a set master key is what makes a missing Authorization header a 401
|
||||
settings: Final = patch("litellm.proxy.proxy_server.general_settings", {"pass_through_endpoints": [yaml_endpoint]}) # test-quality-ok: the method reads this module global; no injection seam
|
||||
yaml_endpoints: Final = patch("litellm.proxy.proxy_server.config_passthrough_endpoints", [yaml_endpoint]) # test-quality-ok: module global holding the YAML endpoints the fix merges in
|
||||
initialize: Final = patch("litellm.proxy.proxy_server.initialize_pass_through_endpoints", AsyncMock()) # test-quality-ok: route registration needs the FastAPI app; auth is the observable here
|
||||
master_key: Final = patch("litellm.proxy.proxy_server.master_key", "sk-master") # test-quality-ok: a set master key is what makes a missing Authorization header a 401
|
||||
with settings, yaml_endpoints, initialize, master_key:
|
||||
await ProxyConfig()._update_general_settings(db_general_settings={"pass_through_endpoints": [db_endpoint]})
|
||||
|
||||
still_open: Final = await user_api_key_auth(request=request, api_key=None)
|
||||
assert still_open.api_key is None
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def app_routes_restored():
|
||||
routes_before: Final = tuple(app.router.routes)
|
||||
yield
|
||||
app.router.routes[:] = routes_before
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.usefixtures("app_routes_restored")
|
||||
async def test_deleting_the_stored_pass_through_row_takes_the_route_out_of_service():
|
||||
"""A pass-through route the database declared has to stop serving when that row is
|
||||
deleted. The proxy's own registry of live pass-through routes is what decides whether
|
||||
a request is routed upstream or falls through to the auth error, so it has to lose the
|
||||
entry on the reload rather than at the next process restart."""
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
InitPassThroughEndpointHelpers,
|
||||
_registered_pass_through_routes,
|
||||
)
|
||||
from litellm.proxy.proxy_server import ProxyConfig, app
|
||||
|
||||
path: Final = f"/v1/deleted-{uuid.uuid4().hex[:8]}"
|
||||
db_endpoint: Final = {"id": "db-1", "path": path, "target": "https://example.com/post"}
|
||||
prior_routes: Final = list(app.routes)
|
||||
prior_registry: Final = dict(_registered_pass_through_routes)
|
||||
|
||||
def live_routes() -> set[str]:
|
||||
return {
|
||||
route for route in InitPassThroughEndpointHelpers.get_all_registered_pass_through_routes() if path in route
|
||||
}
|
||||
|
||||
settings: Final = patch(
|
||||
"litellm.proxy.proxy_server.general_settings", {}
|
||||
) # test-quality-ok: the method reads this module global; no injection seam
|
||||
yaml_endpoints: Final = patch(
|
||||
"litellm.proxy.proxy_server.config_passthrough_endpoints", None
|
||||
) # test-quality-ok: module global holding the YAML endpoints; this case has none
|
||||
app_routes: Final = patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.SafeRouteAdder.add_api_route_if_not_exists"
|
||||
) # test-quality-ok: the registry is the observable; a real route would stay on the shared FastAPI app for the rest of the xdist worker
|
||||
try:
|
||||
with settings, yaml_endpoints, app_routes:
|
||||
pc = ProxyConfig()
|
||||
await pc._update_general_settings(db_general_settings={"pass_through_endpoints": [db_endpoint]})
|
||||
assert live_routes(), "the stored endpoint should be serving before the row is deleted"
|
||||
|
||||
await pc._update_general_settings(db_general_settings={})
|
||||
|
||||
assert live_routes() == set()
|
||||
finally:
|
||||
app.routes[:] = prior_routes
|
||||
_registered_pass_through_routes.clear()
|
||||
_registered_pass_through_routes.update(prior_registry)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.usefixtures("app_routes_restored")
|
||||
async def test_a_stored_pass_through_row_never_disturbs_the_config_declared_routes():
|
||||
"""``pass_through_endpoints`` is config-owned once the file declares it, so writing and then
|
||||
deleting a stored row resolves to the same list both times and the config file's routes keep
|
||||
serving untouched. The stored entry never gets a route of its own."""
|
||||
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
|
||||
InitPassThroughEndpointHelpers,
|
||||
_registered_pass_through_routes,
|
||||
initialize_pass_through_endpoints,
|
||||
)
|
||||
from litellm.proxy.proxy_server import ProxyConfig, app
|
||||
|
||||
marker: Final = uuid.uuid4().hex[:8]
|
||||
config_path: Final = f"/v1/kept-{marker}"
|
||||
db_path: Final = f"/v1/ignored-{marker}"
|
||||
config_endpoint: Final = {"id": f"cfg-{marker}", "path": config_path, "target": "https://example.com/post"}
|
||||
db_endpoint: Final = {"id": f"db-{marker}", "path": db_path, "target": "https://example.com/post"}
|
||||
prior_routes: Final = list(app.routes)
|
||||
prior_registry: Final = dict(_registered_pass_through_routes)
|
||||
|
||||
def live_paths() -> set[str]:
|
||||
registered: Final = InitPassThroughEndpointHelpers.get_all_registered_pass_through_routes()
|
||||
return {path for path in (config_path, db_path) if any(path in route for route in registered)}
|
||||
|
||||
settings: Final = patch(
|
||||
"litellm.proxy.proxy_server.general_settings", {"pass_through_endpoints": [config_endpoint]}
|
||||
) # test-quality-ok: the method reads this module global; no injection seam
|
||||
yaml_endpoints: Final = patch(
|
||||
"litellm.proxy.proxy_server.config_passthrough_endpoints", [config_endpoint]
|
||||
) # test-quality-ok: module global holding the YAML endpoints the reload merges in
|
||||
app_routes: Final = patch(
|
||||
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.SafeRouteAdder.add_api_route_if_not_exists"
|
||||
) # test-quality-ok: the registry is the observable; a real route would stay on the shared FastAPI app for the rest of the xdist worker
|
||||
try:
|
||||
with settings, yaml_endpoints, app_routes:
|
||||
await initialize_pass_through_endpoints(pass_through_endpoints=[config_endpoint])
|
||||
assert live_paths() == {config_path}
|
||||
|
||||
pc = ProxyConfig()
|
||||
await pc._update_general_settings(db_general_settings={"pass_through_endpoints": [db_endpoint]})
|
||||
assert live_paths() == {config_path}
|
||||
|
||||
await pc._update_general_settings(db_general_settings={})
|
||||
|
||||
assert live_paths() == {config_path}
|
||||
finally:
|
||||
app.routes[:] = prior_routes
|
||||
_registered_pass_through_routes.clear()
|
||||
_registered_pass_through_routes.update(prior_registry)
|
||||
with pytest.raises(ProxyException) as locked_down:
|
||||
await user_api_key_auth(request=request, api_key=None)
|
||||
assert locked_down.value.code == "401"
|
||||
|
||||
|
||||
def _fill_user_api_key_cache(cache: DualCache, count: int) -> None:
|
||||
|
|
@ -12442,6 +12325,7 @@ def _config_field_info_client(monkeypatch, user_role):
|
|||
mock_config_table.find_first = AsyncMock(return_value=db_record)
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
|
||||
mock_prisma.writer_db = mock_prisma.db
|
||||
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
|
||||
|
||||
settings = SettingsStore("general_settings")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue