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
|
Schedule advance happens server-side: recurring rows get a fresh
|
||||||
next_run_at, fire_once rows flip to 'fired'.
|
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()
|
prisma_client = _get_prisma_client()
|
||||||
|
|
||||||
parsed_actions: Optional[List[str]] = None
|
parsed_actions: Optional[List[str]] = None
|
||||||
|
|
@ -199,6 +199,7 @@ async def get_due_tasks(
|
||||||
|
|
||||||
rows = await store.claim_due(
|
rows = await store.claim_due(
|
||||||
prisma_client,
|
prisma_client,
|
||||||
|
owner_token=owner_token,
|
||||||
agent_id=user_api_key_dict.agent_id,
|
agent_id=user_api_key_dict.agent_id,
|
||||||
actions=parsed_actions,
|
actions=parsed_actions,
|
||||||
limit=limit,
|
limit=limit,
|
||||||
|
|
|
||||||
|
|
@ -264,6 +264,7 @@ async def cancel_task_for_owner(
|
||||||
async def claim_due(
|
async def claim_due(
|
||||||
prisma_client: Any,
|
prisma_client: Any,
|
||||||
*,
|
*,
|
||||||
|
owner_token: str,
|
||||||
agent_id: Optional[str],
|
agent_id: Optional[str],
|
||||||
actions: Optional[List[str]],
|
actions: Optional[List[str]],
|
||||||
limit: int,
|
limit: int,
|
||||||
|
|
@ -274,7 +275,15 @@ async def claim_due(
|
||||||
APPROVED RAW-SQL EXCEPTION (CLAUDE.md): SELECT FOR UPDATE SKIP LOCKED
|
APPROVED RAW-SQL EXCEPTION (CLAUDE.md): SELECT FOR UPDATE SKIP LOCKED
|
||||||
is required for multi-pod safety and is not expressible via Prisma
|
is required for multi-pod safety and is not expressible via Prisma
|
||||||
model methods. Per-row writes use Prisma inside the same transaction.
|
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)
|
now = datetime.now(timezone.utc)
|
||||||
|
|
||||||
async with prisma_client.db.tx(timeout=timedelta(seconds=30)) as tx:
|
async with prisma_client.db.tx(timeout=timedelta(seconds=30)) as tx:
|
||||||
|
|
@ -287,12 +296,14 @@ async def claim_due(
|
||||||
FROM "LiteLLM_ScheduledTaskTable"
|
FROM "LiteLLM_ScheduledTaskTable"
|
||||||
WHERE status = 'pending'
|
WHERE status = 'pending'
|
||||||
AND next_run_at <= now()
|
AND next_run_at <= now()
|
||||||
AND ($1::text IS NULL OR agent_id = $1)
|
AND owner_token = $1
|
||||||
AND ($2::text[] IS NULL OR action = ANY($2))
|
AND ($2::text IS NULL OR agent_id = $2)
|
||||||
|
AND ($3::text[] IS NULL OR action = ANY($3))
|
||||||
ORDER BY next_run_at
|
ORDER BY next_run_at
|
||||||
LIMIT $3
|
LIMIT $4
|
||||||
FOR UPDATE SKIP LOCKED
|
FOR UPDATE SKIP LOCKED
|
||||||
""",
|
""",
|
||||||
|
owner_token,
|
||||||
agent_id,
|
agent_id,
|
||||||
actions,
|
actions,
|
||||||
limit,
|
limit,
|
||||||
|
|
|
||||||
|
|
@ -193,9 +193,9 @@ class _FakeTx:
|
||||||
|
|
||||||
async def query_raw(self, query: str, *args) -> List[Dict[str, Any]]:
|
async def query_raw(self, query: str, *args) -> List[Dict[str, Any]]:
|
||||||
# Model the production claim query: rows where status='pending',
|
# Model the production claim query: rows where status='pending',
|
||||||
# next_run_at <= now, optional agent_id and actions filters, ordered
|
# next_run_at <= now, owner_token = $1, optional agent_id and
|
||||||
# by next_run_at, limited.
|
# actions filters, ordered by next_run_at, limited.
|
||||||
agent_id, actions, limit = args
|
owner_token, agent_id, actions, limit = args
|
||||||
now = _now()
|
now = _now()
|
||||||
out: List[Dict[str, Any]] = []
|
out: List[Dict[str, Any]] = []
|
||||||
for r in sorted(self._table.rows, key=lambda x: x.next_run_at):
|
for r in sorted(self._table.rows, key=lambda x: x.next_run_at):
|
||||||
|
|
@ -203,6 +203,8 @@ class _FakeTx:
|
||||||
continue
|
continue
|
||||||
if r.next_run_at > now:
|
if r.next_run_at > now:
|
||||||
continue
|
continue
|
||||||
|
if r.owner_token != owner_token:
|
||||||
|
continue
|
||||||
if agent_id is not None and r.agent_id != agent_id:
|
if agent_id is not None and r.agent_id != agent_id:
|
||||||
continue
|
continue
|
||||||
if actions is not None and r.action not in actions:
|
if actions is not None and r.action not in actions:
|
||||||
|
|
@ -697,3 +699,99 @@ class TestLazyExpiry:
|
||||||
body = r.json()
|
body = r.json()
|
||||||
statuses = {t["task_id"]: t["status"] for t in body["tasks"]}
|
statuses = {t["task_id"]: t["status"] for t in body["tasks"]}
|
||||||
assert statuses["exp-1"] == "expired"
|
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