mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix: address Greptile P1 review comments
This commit is contained in:
parent
519095bfe5
commit
e01fe01d35
4 changed files with 98 additions and 14 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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."
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue