Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
44 changes: 31 additions & 13 deletions litellm/proxy/prisma_migration.py
Original file line number Diff line number Diff line change
@@ -1,26 +1,44 @@
# What is this?
## Script to apply initial prisma migration on Docker setup
"""Standalone entrypoint for applying database migrations and generating the Prisma client.

The entrypoint enforces migration failures by default. Set
ENFORCE_PRISMA_MIGRATION_CHECK=false to preserve log-only behavior for migration and
Prisma generate failures.
"""

import os
import subprocess
import sys

sys.path.insert(0, os.path.abspath("./")) # Adds the parent directory to the system path
sys.path.insert(0, os.path.abspath("./"))

from typing import Final

from litellm._logging import verbose_proxy_logger
from litellm.proxy.proxy_cli import run_server
from litellm.secret_managers.main import str_to_bool


def main() -> int:
enforce_prisma_migration_check: Final = str_to_bool(os.getenv("ENFORCE_PRISMA_MIGRATION_CHECK")) is not False
run_server_args: Final = (
("--skip_server_startup", "--enforce_prisma_migration_check")
if enforce_prisma_migration_check
else ("--skip_server_startup",)
)
run_server(run_server_args, standalone_mode=False)

verbose_proxy_logger.info("Running 'prisma generate'...")
result: Final = subprocess.run(("prisma", "generate"), capture_output=True, text=True)
verbose_proxy_logger.info("'prisma generate' stdout: %s", result.stdout)
exit_code: Final = result.returncode

# Call the Click command with standalone_mode=False
run_server(["--skip_server_startup"], standalone_mode=False)
if exit_code != 0:
verbose_proxy_logger.info("'prisma generate' failed with exit code %s.", exit_code)
verbose_proxy_logger.error("'prisma generate' stderr: %s", result.stderr)
if enforce_prisma_migration_check:
return exit_code
return 0

# run prisma generate
verbose_proxy_logger.info("Running 'prisma generate'...")
result: Final = subprocess.run(["prisma", "generate"], capture_output=True, text=True)
verbose_proxy_logger.info("'prisma generate' stdout: %s", result.stdout) # Log stdout
exit_code: Final = result.returncode

if exit_code != 0:
verbose_proxy_logger.info("'prisma generate' failed with exit code %s.", exit_code)
verbose_proxy_logger.error("'prisma generate' stderr: %s", result.stderr) # Log stderr
if __name__ == "__main__":
sys.exit(main())
68 changes: 68 additions & 0 deletions tests/test_litellm/proxy/test_prisma_migration.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,68 @@
import os
from unittest.mock import MagicMock, patch

import pytest

from litellm.proxy import prisma_migration


class TestPrismaMigration:
@patch("litellm.proxy.prisma_migration.subprocess.run")
@patch("litellm.proxy.prisma_migration.run_server")
def test_main_enforces_migration_check_by_default(
self, mock_run_server: MagicMock, mock_subprocess_run: MagicMock
) -> None:
mock_subprocess_run.return_value = MagicMock(returncode=0, stdout="", stderr="")

with patch.dict(os.environ, {}, clear=True):
assert prisma_migration.main() == 0

mock_run_server.assert_called_once_with(
("--skip_server_startup", "--enforce_prisma_migration_check"),
standalone_mode=False,
)

@patch("litellm.proxy.prisma_migration.subprocess.run")
@patch("litellm.proxy.prisma_migration.run_server")
def test_main_disables_migration_check_when_explicitly_false(
self, mock_run_server: MagicMock, mock_subprocess_run: MagicMock
) -> None:
mock_subprocess_run.return_value = MagicMock(returncode=0, stdout="", stderr="")

with patch.dict(os.environ, {"ENFORCE_PRISMA_MIGRATION_CHECK": "false"}, clear=True):
assert prisma_migration.main() == 0

mock_run_server.assert_called_once_with(("--skip_server_startup",), standalone_mode=False)

@patch("litellm.proxy.prisma_migration.subprocess.run")
@patch("litellm.proxy.prisma_migration.run_server")
def test_main_returns_prisma_generate_exit_code_when_enforced(
self, mock_run_server: MagicMock, mock_subprocess_run: MagicMock
) -> None:
mock_subprocess_run.return_value = MagicMock(returncode=7, stdout="", stderr="")

with patch.dict(os.environ, {}, clear=True):
assert prisma_migration.main() == 7

@patch("litellm.proxy.prisma_migration.subprocess.run")
@patch("litellm.proxy.prisma_migration.run_server")
def test_main_ignores_prisma_generate_exit_code_when_disabled(
self, mock_run_server: MagicMock, mock_subprocess_run: MagicMock
) -> None:
mock_subprocess_run.return_value = MagicMock(returncode=7, stdout="", stderr="")

with patch.dict(os.environ, {"ENFORCE_PRISMA_MIGRATION_CHECK": "false"}, clear=True):
assert prisma_migration.main() == 0

@patch("litellm.proxy.prisma_migration.subprocess.run")
@patch("litellm.proxy.prisma_migration.run_server")
def test_main_propagates_migration_failure(
self, mock_run_server: MagicMock, mock_subprocess_run: MagicMock
) -> None:
mock_run_server.side_effect = SystemExit(1)

with patch.dict(os.environ, {}, clear=True):
with pytest.raises(SystemExit, match="1"):
prisma_migration.main()

mock_subprocess_run.assert_not_called()
Loading