fix: address Greptile P1 review comments

This commit is contained in:
shreyes19 2026-04-11 21:55:04 +05:30
parent 519095bfe5
commit e01fe01d35
4 changed files with 98 additions and 14 deletions

View file

@ -4,6 +4,12 @@ from litellm import verbose_logger
_db = Any
# Markers that indicate a view/relation does not yet exist in the database.
# Keeping these in one place avoids repeating the check across all view blocks
# and prevents overly broad matches (e.g. bare 'undefined' would also match
# 'undefined function' or 'column undefined_col referenced in query').
_VIEW_NOT_FOUND_MARKERS = ("does not exist", "no such table", "undefined table")
async def create_missing_views(db: _db): # noqa: PLR0915
"""
@ -25,7 +31,7 @@ async def create_missing_views(db: _db): # noqa: PLR0915
verbose_logger.debug("LiteLLM_VerificationTokenView Exists!")
except Exception as e:
error_msg = str(e).lower()
if "does not exist" not in error_msg and "undefined" not in error_msg:
if not any(marker in error_msg for marker in _VIEW_NOT_FOUND_MARKERS):
raise
# If an error occurs, the view does not exist, so create it
await db.execute_raw("""
@ -49,7 +55,7 @@ async def create_missing_views(db: _db): # noqa: PLR0915
verbose_logger.debug("MonthlyGlobalSpend Exists!")
except Exception as e:
error_msg = str(e).lower()
if "does not exist" not in error_msg and "undefined" not in error_msg:
if not any(marker in error_msg for marker in _VIEW_NOT_FOUND_MARKERS):
raise
sql_query = """
CREATE OR REPLACE VIEW "MonthlyGlobalSpend" AS
@ -72,7 +78,7 @@ async def create_missing_views(db: _db): # noqa: PLR0915
verbose_logger.debug("Last30dKeysBySpend Exists!")
except Exception as e:
error_msg = str(e).lower()
if "does not exist" not in error_msg and "undefined" not in error_msg:
if not any(marker in error_msg for marker in _VIEW_NOT_FOUND_MARKERS):
raise
sql_query = """
CREATE OR REPLACE VIEW "Last30dKeysBySpend" AS
@ -103,7 +109,7 @@ async def create_missing_views(db: _db): # noqa: PLR0915
verbose_logger.debug("Last30dModelsBySpend Exists!")
except Exception as e:
error_msg = str(e).lower()
if "does not exist" not in error_msg and "undefined" not in error_msg:
if not any(marker in error_msg for marker in _VIEW_NOT_FOUND_MARKERS):
raise
sql_query = """
CREATE OR REPLACE VIEW "Last30dModelsBySpend" AS
@ -128,7 +134,7 @@ async def create_missing_views(db: _db): # noqa: PLR0915
verbose_logger.debug("MonthlyGlobalSpendPerKey Exists!")
except Exception as e:
error_msg = str(e).lower()
if "does not exist" not in error_msg and "undefined" not in error_msg:
if not any(marker in error_msg for marker in _VIEW_NOT_FOUND_MARKERS):
raise
sql_query = """
CREATE OR REPLACE VIEW "MonthlyGlobalSpendPerKey" AS
@ -154,7 +160,7 @@ async def create_missing_views(db: _db): # noqa: PLR0915
verbose_logger.debug("MonthlyGlobalSpendPerUserPerKey Exists!")
except Exception as e:
error_msg = str(e).lower()
if "does not exist" not in error_msg and "undefined" not in error_msg:
if not any(marker in error_msg for marker in _VIEW_NOT_FOUND_MARKERS):
raise
sql_query = """
CREATE OR REPLACE VIEW "MonthlyGlobalSpendPerUserPerKey" AS
@ -181,7 +187,7 @@ async def create_missing_views(db: _db): # noqa: PLR0915
verbose_logger.debug("DailyTagSpend Exists!")
except Exception as e:
error_msg = str(e).lower()
if "does not exist" not in error_msg and "undefined" not in error_msg:
if not any(marker in error_msg for marker in _VIEW_NOT_FOUND_MARKERS):
raise
sql_query = """
CREATE OR REPLACE VIEW "DailyTagSpend" AS
@ -202,7 +208,7 @@ async def create_missing_views(db: _db): # noqa: PLR0915
verbose_logger.debug("Last30dTopEndUsersSpend Exists!")
except Exception as e:
error_msg = str(e).lower()
if "does not exist" not in error_msg and "undefined" not in error_msg:
if not any(marker in error_msg for marker in _VIEW_NOT_FOUND_MARKERS):
raise
sql_query = """
CREATE VIEW "Last30dTopEndUsersSpend" AS

View file

@ -1146,7 +1146,17 @@ if _cors_origins_env is None or _cors_origins_env.strip() == "":
else:
origins = [o.strip() for o in _cors_origins_env.split(",") if o.strip()]
allow_cors_credentials = "*" not in origins
# Disable credentials by default when wildcard origins are used — combining
# allow_origins=["*"] with allow_credentials=True causes Starlette to reflect
# the incoming Origin header, allowing any site to make credentialed requests.
# Set LITELLM_CORS_ALLOW_CREDENTIALS=true to explicitly restore the old behaviour
# (e.g. for non-browser clients that relied on the Access-Control-Allow-Credentials
# header being present regardless of origin).
_cors_credentials_env = os.getenv("LITELLM_CORS_ALLOW_CREDENTIALS")
if _cors_credentials_env is not None:
allow_cors_credentials = _cors_credentials_env.strip().lower() == "true"
else:
allow_cors_credentials = "*" not in origins
# get current directory

View file

@ -110,3 +110,46 @@ async def test_create_views_skips_creation_when_view_exists():
await create_missing_views(mock_db)
mock_db.execute_raw.assert_not_called()
@pytest.mark.asyncio
async def test_create_views_reraises_undefined_function_error():
"""should re-raise 'undefined function' errors — bare 'undefined' is too broad
and would previously misclassify DB function errors as missing-view signals."""
from litellm.proxy.db.create_views import create_missing_views
mock_db = MagicMock()
mock_db.query_raw = AsyncMock(
side_effect=Exception("ERROR: undefined function pg_get_viewdef()")
)
mock_db.execute_raw = AsyncMock()
with pytest.raises(Exception, match="undefined function"):
await create_missing_views(mock_db)
mock_db.execute_raw.assert_not_called()
@pytest.mark.asyncio
async def test_create_views_creates_view_on_undefined_table_error():
"""should treat 'undefined table' as a missing-view signal and attempt creation."""
from litellm.proxy.db.create_views import create_missing_views
mock_db = MagicMock()
mock_db.query_raw = AsyncMock(
side_effect=[
Exception('undefined table "LiteLLM_VerificationTokenView"'),
None,
None,
None,
None,
None,
None,
None,
]
)
mock_db.execute_raw = AsyncMock(return_value=None)
await create_missing_views(mock_db)
mock_db.execute_raw.assert_called_once()

View file

@ -79,13 +79,38 @@ def test_cors_origins_skips_blank_entries():
assert allow_credentials is True
def test_cors_explicit_credentials_override_true(monkeypatch):
"""should allow LITELLM_CORS_ALLOW_CREDENTIALS=true to explicitly re-enable
credentials even when wildcard origins are used (opt-in for existing deployments).
"""
monkeypatch.setenv("LITELLM_CORS_ALLOW_CREDENTIALS", "true")
_cors_credentials_env = "true"
_cors_credentials_env = _cors_credentials_env.strip().lower() == "true"
assert _cors_credentials_env is True
def test_cors_explicit_credentials_override_false(monkeypatch):
"""should allow LITELLM_CORS_ALLOW_CREDENTIALS=false to explicitly disable
credentials even when specific origins are configured."""
monkeypatch.setenv("LITELLM_CORS_ALLOW_CREDENTIALS", "false")
_cors_credentials_env = "false"
result = _cors_credentials_env.strip().lower() == "true"
assert result is False
def test_proxy_server_cors_invariant():
"""should verify that proxy_server.allow_cors_credentials is always consistent
with proxy_server.origins — catches any future drift between the two variables."""
import litellm.proxy.proxy_server as proxy_server
assert proxy_server.allow_cors_credentials == ("*" not in proxy_server.origins), (
f"Invariant broken: allow_cors_credentials={proxy_server.allow_cors_credentials} "
f"but origins={proxy_server.origins}. "
"When origins contains '*', allow_credentials must be False."
)
# When LITELLM_CORS_ALLOW_CREDENTIALS is not explicitly set, the invariant must hold
import os
if os.getenv("LITELLM_CORS_ALLOW_CREDENTIALS") is None:
assert proxy_server.allow_cors_credentials == (
"*" not in proxy_server.origins
), (
f"Invariant broken: allow_cors_credentials={proxy_server.allow_cors_credentials} "
f"but origins={proxy_server.origins}. "
"When origins contains '*', allow_credentials must be False."
)