mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-09-25 14:06:18 +00:00
fix(gateway): scope edit/regenerate helper fallbacks by data identity (#5483)
The message edit/regenerate prepare chain resolves source runs through _resolve_run_id_for_message, _require_successful_source_run and _find_interrupted_target_run_id, which filtered by the authorization identity (get_current_user) — the same conflation fixed for the runs and messages read endpoints in #5448, left out of scope there. For trusted internal callers the authorization identity never matches the raw owner-stamped run rows, so regenerate/edit-regenerate prepare fails with 409 on threads the caller is authorized on (#5482). The three helpers now resolve their filter id through _run_scope_user_id as well; both prepare endpoints keep their owner_check=True authorization and browser/API sessions keep the per-user filter. Regression tests extend test_thread_runs_internal_scope.py to the helper fallback paths: internal callers resolve raw-owner-stamped runs (including the interrupted-run and status-409 paths), browser sessions keep the per-user filter and 409 on cross-user runs.
This commit is contained in:
parent
a022be195a
commit
e493390aea
@ -590,7 +590,7 @@ async def _find_target_run_id(
|
|||||||
return source_run_id
|
return source_run_id
|
||||||
|
|
||||||
run_mgr = get_run_manager(request)
|
run_mgr = get_run_manager(request)
|
||||||
user_id = await get_current_user(request)
|
user_id = await _run_scope_user_id(request)
|
||||||
records = await run_mgr.list_by_thread(thread_id, user_id=user_id, limit=10)
|
records = await run_mgr.list_by_thread(thread_id, user_id=user_id, limit=10)
|
||||||
fallback_record = next(
|
fallback_record = next(
|
||||||
(record for record in records if record.status == RunStatus.success and _run_last_ai_matches_message(record, target_message)),
|
(record for record in records if record.status == RunStatus.success and _run_last_ai_matches_message(record, target_message)),
|
||||||
@ -677,7 +677,7 @@ def _run_status_value(record: Any) -> str | None:
|
|||||||
|
|
||||||
async def _require_successful_source_run(thread_id: str, run_id: str, request: Request) -> RunRecord:
|
async def _require_successful_source_run(thread_id: str, run_id: str, request: Request) -> RunRecord:
|
||||||
run_mgr = get_run_manager(request)
|
run_mgr = get_run_manager(request)
|
||||||
user_id = await get_current_user(request)
|
user_id = await _run_scope_user_id(request)
|
||||||
record = await run_mgr.get(run_id, user_id=user_id)
|
record = await run_mgr.get(run_id, user_id=user_id)
|
||||||
if record is None:
|
if record is None:
|
||||||
# The run-event journal is the authoritative lookup above. This fallback
|
# The run-event journal is the authoritative lookup above. This fallback
|
||||||
@ -704,7 +704,7 @@ async def _find_interrupted_target_run_id(
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
run_mgr = get_run_manager(request)
|
run_mgr = get_run_manager(request)
|
||||||
user_id = await get_current_user(request)
|
user_id = await _run_scope_user_id(request)
|
||||||
record = await run_mgr.get(source_run_id, user_id=user_id)
|
record = await run_mgr.get(source_run_id, user_id=user_id)
|
||||||
if record is None:
|
if record is None:
|
||||||
records = await run_mgr.list_by_thread(thread_id, user_id=user_id, limit=20)
|
records = await run_mgr.list_by_thread(thread_id, user_id=user_id, limit=20)
|
||||||
|
|||||||
@ -13,10 +13,11 @@ browser/API sessions.
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
from types import SimpleNamespace
|
||||||
from uuid import UUID
|
from uuid import UUID
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from fastapi import FastAPI, Request, Response
|
from fastapi import FastAPI, HTTPException, Request, Response
|
||||||
from fastapi.testclient import TestClient
|
from fastapi.testclient import TestClient
|
||||||
from starlette.middleware.base import BaseHTTPMiddleware
|
from starlette.middleware.base import BaseHTTPMiddleware
|
||||||
|
|
||||||
@ -116,8 +117,8 @@ def _internal_user(owner_raw: str | None):
|
|||||||
return get_internal_user(owner_user_id=owner_raw)
|
return get_internal_user(owner_user_id=owner_raw)
|
||||||
|
|
||||||
|
|
||||||
def _seed_run(store: MemoryRunStore, run_id: str, *, user_id: str | None) -> None:
|
def _seed_run(store: MemoryRunStore, run_id: str, *, user_id: str | None, status: str = "success") -> None:
|
||||||
asyncio.run(store.put(run_id, thread_id=THREAD_ID, user_id=user_id, status="success"))
|
asyncio.run(store.put(run_id, thread_id=THREAD_ID, user_id=user_id, status=status))
|
||||||
|
|
||||||
|
|
||||||
def _seed_message(event_store: MemoryRunEventStore, run_id: str, message_id: str) -> None:
|
def _seed_message(event_store: MemoryRunEventStore, run_id: str, message_id: str) -> None:
|
||||||
@ -315,3 +316,83 @@ def test_browser_session_messages_keep_per_user_filter(mixed_owner_store: Memory
|
|||||||
assert {row["content"]["id"] for row in response.json()} == {"msg-owner", "msg-browser"}
|
assert {row["content"]["id"] for row in response.json()} == {"msg-owner", "msg-browser"}
|
||||||
assert run_store.list_by_thread_user_ids and all(uid == str(BROWSER_USER_ID) for uid in run_store.list_by_thread_user_ids)
|
assert run_store.list_by_thread_user_ids and all(uid == str(BROWSER_USER_ID) for uid in run_store.list_by_thread_user_ids)
|
||||||
assert feedback_repo.list_by_thread_user_ids == [str(BROWSER_USER_ID)]
|
assert feedback_repo.list_by_thread_user_ids == [str(BROWSER_USER_ID)]
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# edit/regenerate helper fallback paths (#5482)
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _helper_request(*, user, auth_source: str, run_store, event_store):
|
||||||
|
"""Minimal Request stand-in: the helpers only touch state and app.state."""
|
||||||
|
app_state = SimpleNamespace(run_manager=RunManager(store=run_store), run_event_store=event_store)
|
||||||
|
return SimpleNamespace(
|
||||||
|
state=SimpleNamespace(user=user, auth_source=auth_source),
|
||||||
|
app=SimpleNamespace(state=app_state),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_helper_fallback_paths_resolve_internal_caller_runs() -> None:
|
||||||
|
"""The edit/regenerate helper fallbacks must use the data identity (#5482)."""
|
||||||
|
store = _RecordingRunStore()
|
||||||
|
_seed_run(store, RUN_OWNER, user_id=OWNER_RAW, status="interrupted")
|
||||||
|
_seed_run(store, RUN_BROWSER, user_id=str(BROWSER_USER_ID), status="success")
|
||||||
|
request = _helper_request(
|
||||||
|
user=_internal_user(OWNER_RAW),
|
||||||
|
auth_source=AUTH_SOURCE_INTERNAL,
|
||||||
|
run_store=store,
|
||||||
|
event_store=MemoryRunEventStore(),
|
||||||
|
)
|
||||||
|
|
||||||
|
# The owner-stamped interrupted run resolves through the raw owner stamp.
|
||||||
|
interrupted = asyncio.run(thread_runs._find_interrupted_target_run_id(THREAD_ID, {"additional_kwargs": {"run_id": RUN_OWNER}}, request))
|
||||||
|
assert interrupted == RUN_OWNER
|
||||||
|
assert store.get_user_ids[-1] is None
|
||||||
|
|
||||||
|
# An interrupted run is not an editable source run, but the lookup itself
|
||||||
|
# must have reached it (409 for status, not for a missing record).
|
||||||
|
with pytest.raises(HTTPException) as exc:
|
||||||
|
asyncio.run(thread_runs._require_successful_source_run(THREAD_ID, RUN_OWNER, request))
|
||||||
|
assert exc.value.status_code == 409
|
||||||
|
assert "successful" in exc.value.detail
|
||||||
|
assert store.get_user_ids[-1] is None
|
||||||
|
|
||||||
|
# Fallback scan without any event-store or kwargs anchors still scans the
|
||||||
|
# authorized thread unfiltered (and 409s on the miss).
|
||||||
|
with pytest.raises(HTTPException) as exc2:
|
||||||
|
asyncio.run(
|
||||||
|
thread_runs._find_target_run_id(
|
||||||
|
THREAD_ID,
|
||||||
|
"missing-message",
|
||||||
|
{"content": "unmatched"},
|
||||||
|
{"additional_kwargs": {}},
|
||||||
|
request,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
assert exc2.value.status_code == 409
|
||||||
|
assert store.list_by_thread_user_ids and store.list_by_thread_user_ids[-1] is None
|
||||||
|
|
||||||
|
|
||||||
|
def test_helper_fallback_paths_keep_per_user_filter_for_browser_sessions() -> None:
|
||||||
|
"""Browser sessions keep the per-user filter in the helper fallbacks."""
|
||||||
|
store = _RecordingRunStore()
|
||||||
|
_seed_run(store, RUN_OWNER, user_id=OWNER_RAW, status="interrupted")
|
||||||
|
_seed_run(store, RUN_BROWSER, user_id=str(BROWSER_USER_ID), status="success")
|
||||||
|
request = _helper_request(
|
||||||
|
user=_browser_user(),
|
||||||
|
auth_source=AUTH_SOURCE_SESSION,
|
||||||
|
run_store=store,
|
||||||
|
event_store=MemoryRunEventStore(),
|
||||||
|
)
|
||||||
|
|
||||||
|
# The owner-stamped run is invisible under the browser user's filter.
|
||||||
|
assert asyncio.run(thread_runs._find_interrupted_target_run_id(THREAD_ID, {"additional_kwargs": {"run_id": RUN_OWNER}}, request)) is None
|
||||||
|
assert store.get_user_ids[-1] == str(BROWSER_USER_ID)
|
||||||
|
|
||||||
|
# Their own successful run still resolves, and a cross-user one 409s.
|
||||||
|
record = asyncio.run(thread_runs._require_successful_source_run(THREAD_ID, RUN_BROWSER, request))
|
||||||
|
assert record.run_id == RUN_BROWSER
|
||||||
|
with pytest.raises(HTTPException) as exc:
|
||||||
|
asyncio.run(thread_runs._require_successful_source_run(THREAD_ID, RUN_OWNER, request))
|
||||||
|
assert exc.value.status_code == 409
|
||||||
|
assert store.get_user_ids[-1] == str(BROWSER_USER_ID)
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user