mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix /v1/tasks/due cross-tenant claim bypass
Greptile/Veria flagged a 7/10 in claim_due:
WHERE status = 'pending'
AND ($1::text IS NULL OR agent_id = $1)
AND ($2::text[] IS NULL OR action = ANY($2))
Most LiteLLM keys have no agent_id set, so $1 was NULL and the IS NULL
branch made the predicate always true. Any authenticated key with NULL
agent_id could call GET /v1/tasks/due and receive every pending task in
the table — including check_prompt and action_args content — and have
those rows mutated (next_run_at advanced or status flipped to 'fired').
Fix: claim_due now requires owner_token (the calling key's hashed token)
and the SQL filters on it unconditionally. agent_id remains an optional
metadata filter applied on top, never the authorization scope.
The /v1/tasks/due endpoint passes user_api_key_dict.token through to
claim_due. Caller never supplies owner_token from the request body — same
contract as create/list/get/update/delete/report.
Regression tests in TestDueTenantIsolation cover both branches:
- caller with NULL agent_id does NOT see foreign rows, and foreign
rows are not mutated by the bypass attempt;
- caller with NULL agent_id still claims its own rows.
Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
parent
8587f3b770
commit
4a783242f2
3 changed files with 117 additions and 7 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue