fix(proxy): resolve db generate schema from litellm package, not proxy-extras

Addresses review feedback (jamesmyatt): db generate now resolves schema.prisma
via PrismaManager (litellm/proxy/schema.prisma) — the same schema db push uses —
instead of reaching into the sibling litellm-proxy-extras package. Keeps the
generate path consistent with the existing push path and corrects the docstring,
which wrongly claimed db push uses the proxy-extras schema. proxy-extras is still
imported for prisma-binary discovery and PATH injection (not the schema).
This commit is contained in:
Stephen Chin 2026-06-17 07:34:53 -07:00
parent c4a0792c99
commit 83e6c28531
2 changed files with 22 additions and 18 deletions

View file

@ -7,11 +7,16 @@ import sysconfig
import click
# Import Prisma helpers from litellm-proxy-extras at module level.
# These are optional dependencies; the ImportError is handled at call time.
# The Prisma schema ships inside the litellm package itself (litellm/proxy/
# schema.prisma); PrismaManager resolves that directory. This is the same
# schema `db push` operates on, so `db generate` stays consistent with it.
from litellm.proxy.db.prisma_client import PrismaManager
# The prisma *binary* discovery and PATH-injection helpers live in
# litellm-proxy-extras. These are optional dependencies; the ImportError is
# handled at call time so `litellm` works without the proxy extras installed.
try:
from litellm_proxy_extras.utils import (
ProxyExtrasDBManager,
_get_prisma_command,
_get_prisma_env,
)
@ -19,7 +24,6 @@ try:
_PROXY_EXTRAS_AVAILABLE = True
except ImportError:
_PROXY_EXTRAS_AVAILABLE = False
ProxyExtrasDBManager = None # type: ignore[assignment]
_get_prisma_command = None # type: ignore[assignment]
_get_prisma_env = None # type: ignore[assignment]
@ -64,14 +68,14 @@ def db() -> None:
@db.command(name="generate")
def db_generate() -> None:
"""Generate the Prisma client using the schema bundled with litellm-proxy-extras.
"""Generate the Prisma client using litellm's bundled schema.
Runs: prisma generate --schema <path_to_schema.prisma>
I resolve the schema path from the installed litellm-proxy-extras package,
so you never need to know internal site-packages paths. This fixes the gap
where `migrate deploy` and `db push` already use the bundled schema but
there was no equivalent for the generate step.
The schema is resolved from the litellm package (litellm/proxy/schema.prisma)
via PrismaManager, which is the same schema `db push` operates on. This keeps
generate consistent with the existing push path and closes the gap where there
was no out-of-the-box `generate` step after a plain pip install.
"""
if not _PROXY_EXTRAS_AVAILABLE:
click.echo(
@ -81,7 +85,7 @@ def db_generate() -> None:
)
raise SystemExit(1)
prisma_dir = ProxyExtrasDBManager._get_prisma_dir()
prisma_dir = PrismaManager._get_prisma_dir()
schema_path = os.path.join(prisma_dir, "schema.prisma")
if not os.path.exists(schema_path):

View file

@ -36,7 +36,7 @@ def test_db_generate_success(cli_runner):
mock_run = MagicMock(return_value=MagicMock(returncode=0))
with (
patch(
"litellm.proxy.client.cli.commands.db.ProxyExtrasDBManager._get_prisma_dir",
"litellm.proxy.client.cli.commands.db.PrismaManager._get_prisma_dir",
return_value="/fake/prisma/dir",
),
patch("litellm.proxy.client.cli.commands.db._get_prisma_command", return_value="prisma"),
@ -60,8 +60,8 @@ def test_db_generate_schema_path_uses_get_prisma_dir(cli_runner):
mock_run = MagicMock(return_value=MagicMock(returncode=0))
with (
patch(
"litellm.proxy.client.cli.commands.db.ProxyExtrasDBManager._get_prisma_dir",
return_value="/custom/extras/dir",
"litellm.proxy.client.cli.commands.db.PrismaManager._get_prisma_dir",
return_value="/custom/litellm/proxy/dir",
),
patch("litellm.proxy.client.cli.commands.db._get_prisma_command", return_value="prisma"),
patch("litellm.proxy.client.cli.commands.db._get_prisma_env", return_value=None),
@ -73,14 +73,14 @@ def test_db_generate_schema_path_uses_get_prisma_dir(cli_runner):
assert result.exit_code == 0, result.output
call_args = mock_run.call_args[0][0]
schema_arg = call_args[call_args.index("--schema") + 1]
assert schema_arg == "/custom/extras/dir/schema.prisma"
assert schema_arg == "/custom/litellm/proxy/dir/schema.prisma"
def test_db_generate_prisma_failure(cli_runner):
mock_run = MagicMock(side_effect=subprocess.CalledProcessError(1, "prisma"))
with (
patch(
"litellm.proxy.client.cli.commands.db.ProxyExtrasDBManager._get_prisma_dir",
"litellm.proxy.client.cli.commands.db.PrismaManager._get_prisma_dir",
return_value="/fake/prisma/dir",
),
patch("litellm.proxy.client.cli.commands.db._get_prisma_command", return_value="prisma"),
@ -97,7 +97,7 @@ def test_db_generate_prisma_failure(cli_runner):
def test_db_generate_schema_missing(cli_runner):
with (
patch(
"litellm.proxy.client.cli.commands.db.ProxyExtrasDBManager._get_prisma_dir",
"litellm.proxy.client.cli.commands.db.PrismaManager._get_prisma_dir",
return_value="/fake/prisma/dir",
),
patch("litellm.proxy.client.cli.commands.db._get_prisma_command", return_value="prisma"),
@ -166,7 +166,7 @@ def test_db_generate_env_includes_scripts_dir_on_path(cli_runner):
with (
patch(
"litellm.proxy.client.cli.commands.db.ProxyExtrasDBManager._get_prisma_dir",
"litellm.proxy.client.cli.commands.db.PrismaManager._get_prisma_dir",
return_value="/fake/prisma/dir",
),
patch(
@ -213,7 +213,7 @@ def test_db_generate_env_does_not_duplicate_scripts_dir(cli_runner):
with (
patch(
"litellm.proxy.client.cli.commands.db.ProxyExtrasDBManager._get_prisma_dir",
"litellm.proxy.client.cli.commands.db.PrismaManager._get_prisma_dir",
return_value="/fake/prisma/dir",
),
patch(