From e01fe01d35d825ab946c1b124bc2517b6ecf9b4d Mon Sep 17 00:00:00 2001 From: shreyes19 Date: Sat, 11 Apr 2026 21:55:04 +0530 Subject: [PATCH] fix: address Greptile P1 review comments --- litellm/proxy/db/create_views.py | 22 ++++++---- litellm/proxy/proxy_server.py | 12 +++++- .../proxy/db/test_create_views.py | 43 +++++++++++++++++++ tests/test_litellm/proxy/test_cors_config.py | 35 ++++++++++++--- 4 files changed, 98 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/db/create_views.py b/litellm/proxy/db/create_views.py index 54eee3c3bcb..2326f495a5f 100644 --- a/litellm/proxy/db/create_views.py +++ b/litellm/proxy/db/create_views.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index ef67781918a..bdb1896b28d 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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 diff --git a/tests/test_litellm/proxy/db/test_create_views.py b/tests/test_litellm/proxy/db/test_create_views.py index aacbcd460ac..1a90b4c204d 100644 --- a/tests/test_litellm/proxy/db/test_create_views.py +++ b/tests/test_litellm/proxy/db/test_create_views.py @@ -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() diff --git a/tests/test_litellm/proxy/test_cors_config.py b/tests/test_litellm/proxy/test_cors_config.py index 4af939e903d..2a63fa656ce 100644 --- a/tests/test_litellm/proxy/test_cors_config.py +++ b/tests/test_litellm/proxy/test_cors_config.py @@ -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." + )