diff --git a/litellm/proxy/scheduled_tasks/endpoints.py b/litellm/proxy/scheduled_tasks/endpoints.py index 628bd8c7f0a..35a77597d88 100644 --- a/litellm/proxy/scheduled_tasks/endpoints.py +++ b/litellm/proxy/scheduled_tasks/endpoints.py @@ -188,7 +188,7 @@ async def get_due_tasks( Schedule advance happens server-side: recurring rows get a fresh next_run_at, fire_once rows flip to 'fired'. """ - _require_token(user_api_key_dict) + owner_token = _require_token(user_api_key_dict) prisma_client = _get_prisma_client() parsed_actions: Optional[List[str]] = None @@ -199,6 +199,7 @@ async def get_due_tasks( rows = await store.claim_due( prisma_client, + owner_token=owner_token, agent_id=user_api_key_dict.agent_id, actions=parsed_actions, limit=limit, diff --git a/litellm/proxy/scheduled_tasks/store.py b/litellm/proxy/scheduled_tasks/store.py index 4c770441d0f..a652ff4923d 100644 --- a/litellm/proxy/scheduled_tasks/store.py +++ b/litellm/proxy/scheduled_tasks/store.py @@ -264,6 +264,7 @@ async def cancel_task_for_owner( async def claim_due( prisma_client: Any, *, + owner_token: str, agent_id: Optional[str], actions: Optional[List[str]], limit: int, @@ -274,7 +275,15 @@ async def claim_due( APPROVED RAW-SQL EXCEPTION (CLAUDE.md): SELECT FOR UPDATE SKIP LOCKED is required for multi-pod safety and is not expressible via Prisma model methods. Per-row writes use Prisma inside the same transaction. + + SECURITY: owner_token is required. Without it, any authenticated key + whose agent_id is NULL would have matched every task in the table via + the `($n::text IS NULL OR agent_id = $n)` pattern, leaking + check_prompt / action_args across tenants. Always scope claims to the + rows the calling key owns. """ + if not owner_token: + raise ValueError("owner_token is required for claim_due") now = datetime.now(timezone.utc) async with prisma_client.db.tx(timeout=timedelta(seconds=30)) as tx: @@ -287,12 +296,14 @@ async def claim_due( FROM "LiteLLM_ScheduledTaskTable" WHERE status = 'pending' AND next_run_at <= now() - AND ($1::text IS NULL OR agent_id = $1) - AND ($2::text[] IS NULL OR action = ANY($2)) + AND owner_token = $1 + AND ($2::text IS NULL OR agent_id = $2) + AND ($3::text[] IS NULL OR action = ANY($3)) ORDER BY next_run_at - LIMIT $3 + LIMIT $4 FOR UPDATE SKIP LOCKED """, + owner_token, agent_id, actions, limit, diff --git a/tests/test_litellm/proxy/scheduled_tasks/test_endpoints.py b/tests/test_litellm/proxy/scheduled_tasks/test_endpoints.py index 5efd0e683cf..3e267c14d3e 100644 --- a/tests/test_litellm/proxy/scheduled_tasks/test_endpoints.py +++ b/tests/test_litellm/proxy/scheduled_tasks/test_endpoints.py @@ -193,9 +193,9 @@ class _FakeTx: async def query_raw(self, query: str, *args) -> List[Dict[str, Any]]: # Model the production claim query: rows where status='pending', - # next_run_at <= now, optional agent_id and actions filters, ordered - # by next_run_at, limited. - agent_id, actions, limit = args + # next_run_at <= now, owner_token = $1, optional agent_id and + # actions filters, ordered by next_run_at, limited. + owner_token, agent_id, actions, limit = args now = _now() out: List[Dict[str, Any]] = [] for r in sorted(self._table.rows, key=lambda x: x.next_run_at): @@ -203,6 +203,8 @@ class _FakeTx: continue if r.next_run_at > now: continue + if r.owner_token != owner_token: + continue if agent_id is not None and r.agent_id != agent_id: continue if actions is not None and r.action not in actions: @@ -697,3 +699,99 @@ class TestLazyExpiry: body = r.json() statuses = {t["task_id"]: t["status"] for t in body["tasks"]} assert statuses["exp-1"] == "expired" + + +class TestDueTenantIsolation: + """ + Regression for the bypass surfaced in PR review (Greptile/Veria): + ($1::text IS NULL OR agent_id = $1) + With agent_id=NULL on the calling key, the IS NULL branch matched + every row in the table and leaked check_prompt / action_args across + tenants. Fix: claim must always scope by owner_token. + """ + + def setup_method(self): + self.prisma = _make_prisma() + self.client = _make_app(self.prisma) + + def test_caller_with_null_agent_id_does_not_see_other_owners_tasks(self): + # Seed two due rows owned by a different key, no agent_id on either. + self.prisma.db.litellm_scheduledtasktable.rows.append( + _make_row( + task_id="other-1", + owner_token="hashed-token-foreign", + agent_id=None, + next_run_at=_now() - timedelta(seconds=10), + check_prompt="leaked secret prompt", + ) + ) + self.prisma.db.litellm_scheduledtasktable.rows.append( + _make_row( + task_id="other-2", + owner_token="hashed-token-foreign", + agent_id=None, + next_run_at=_now() - timedelta(seconds=10), + action_args={"channel": "secret"}, + ) + ) + + # Caller has its own (different) hashed owner_token and NULL + # agent_id (the override below intentionally omits agent_id, but the + # default _make_app fixture sets one — patch the override). + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + def _null_agent_auth(): + return UserAPIKeyAuth( + api_key=_TEST_API_KEY, + user_id="user-a", + team_id="team-a", + # agent_id intentionally None — this is the bypass case. + ) + + self.client.app.dependency_overrides[user_api_key_auth] = _null_agent_auth + + with _patch_prisma(self.prisma): + r = self.client.get("/v1/tasks/due") + assert r.status_code == 200 + body = r.json() + ids = [t["task_id"] for t in body["tasks"]] + assert ( + "other-1" not in ids + ), "TENANT BYPASS: caller with NULL agent_id claimed foreign task other-1" + assert ( + "other-2" not in ids + ), "TENANT BYPASS: caller with NULL agent_id claimed foreign task other-2" + + # Foreign rows must not have been mutated either (no schedule advance, + # no flip-to-fired by the bypass call). + for r in self.prisma.db.litellm_scheduledtasktable.rows: + if r.task_id in ("other-1", "other-2"): + assert ( + r.status == "pending" + ), f"TENANT BYPASS: foreign row {r.task_id} was mutated to {r.status!r}" + + def test_caller_sees_own_tasks_with_null_agent_id(self): + # Sanity: caller with NULL agent_id still claims its own rows. + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + def _null_agent_auth(): + return UserAPIKeyAuth( + api_key=_TEST_API_KEY, + user_id="user-a", + team_id="team-a", + ) + + self.client.app.dependency_overrides[user_api_key_auth] = _null_agent_auth + + self.prisma.db.litellm_scheduledtasktable.rows.append( + _make_row( + task_id="own-1", + agent_id=None, + next_run_at=_now() - timedelta(seconds=10), + ) + ) + with _patch_prisma(self.prisma): + r = self.client.get("/v1/tasks/due") + assert r.status_code == 200 + ids = [t["task_id"] for t in r.json()["tasks"]] + assert ids == ["own-1"]