diff --git a/backend/app/api/routes.py b/backend/app/api/routes.py index 74a80a3..635d316 100644 --- a/backend/app/api/routes.py +++ b/backend/app/api/routes.py @@ -324,7 +324,7 @@ async def ask( qa_service: QAService = Depends(get_qa_service), accept: str = "application/json", ) -> AskResponse | StreamingResponse: - user_id = user.id if user else None + user_id = (user.id or None) if user else None try: if "text/event-stream" in accept: # SSE: generator handles its own recording via try/finally @@ -355,7 +355,7 @@ async def research( service: ResearchService = Depends(get_research_service), accept: str = "application/json", ) -> ResearchResponse | StreamingResponse: - user_id = user.id if user else None + user_id = (user.id or None) if user else None try: if request.research_type == "codemap": if "text/event-stream" in accept: @@ -395,7 +395,7 @@ async def list_wikis( management: WikiManagementService = Depends(get_wiki_management), service: WikiService = Depends(get_wiki_service), ) -> WikiListResponse: - user_id = user.id if user else None + user_id = (user.id or None) if user else None result = await management.list_wikis(user_id=user_id) completed_wiki_ids = {w.wiki_id for w in result.wikis} @@ -415,6 +415,9 @@ async def list_wikis( best_inv = None for inv in service.invocations.values(): if inv.wiki_id == wiki.wiki_id: + # Mirror B2: don't attach another user's invocation metadata + if inv.owner_id and (not user_id or inv.owner_id != user_id): + continue if inv.status in ("generating", "running"): best_inv = inv break @@ -434,8 +437,13 @@ async def list_wikis( wiki.progress = best_inv.progress wiki.error = None - # Add active/failed invocations not yet in completed list + # Add active/failed invocations not yet in completed list. + # Mirror the DB visibility rule: own wikis + legacy unowned (owner_id==""). + # Owned invocations from other users are never shown — the caller would + # get a 404 trying to open them, so leaking them in the list is wrong. for inv in service.invocations.values(): + if inv.owner_id and (not user_id or inv.owner_id != user_id): + continue if inv.wiki_id not in completed_wiki_ids: # Auto-register completed invocations into DB so get_wiki works if inv.status == "complete" and inv.repo_url: @@ -534,7 +542,7 @@ async def get_wiki( service: WikiService = Depends(get_wiki_service), ) -> dict: """Get wiki detail with pages and their content.""" - user_id = user.id if user else None + user_id = (user.id or None) if user else None wiki_record = await management.get_wiki(wiki_id, user_id=user_id) wiki_meta = WikiManagementService._record_to_summary(wiki_record, user_id) if wiki_record else None @@ -555,7 +563,7 @@ async def get_wiki( active_invocation = inv # keep first match as fallback # B2: Don't leak in-flight invocation metadata to non-owners - if active_invocation and user_id and active_invocation.owner_id and active_invocation.owner_id != user_id: + if active_invocation and active_invocation.owner_id and (not user_id or active_invocation.owner_id != user_id): active_invocation = None # If DB has a record in a non-complete state (new: registered at generation start), @@ -749,7 +757,7 @@ async def search_wiki( """Full-text search over wiki pages with graph-expansion re-ranking.""" from app.core.wiki_search_engine import WikiSearchEngine - user_id = user.id if user else None + user_id = (user.id or None) if user else None wiki = await management.get_wiki(wiki_id, user_id=user_id) if wiki is None: raise HTTPException(404, f"Wiki not found: {wiki_id}") @@ -797,7 +805,7 @@ async def list_wiki_pages( index_cache=Depends(get_wiki_index_cache), ) -> WikiPageListResponse: """List all pages in a wiki with their titles and descriptions.""" - user_id = user.id if user else None + user_id = (user.id or None) if user else None wiki = await management.get_wiki(wiki_id, user_id=user_id) if wiki is None: raise HTTPException(404, f"Wiki not found: {wiki_id}") @@ -831,7 +839,7 @@ async def get_page_neighbors( index_cache=Depends(get_wiki_index_cache), ) -> PageNeighborsResponse: """Return the wikilink graph neighborhood for a single wiki page.""" - user_id = user.id if user else None + user_id = (user.id or None) if user else None wiki = await management.get_wiki(wiki_id, user_id=user_id) if wiki is None: raise HTTPException(404, f"Wiki not found: {wiki_id}") @@ -861,7 +869,7 @@ async def get_wiki_page_by_title( index_cache=Depends(get_wiki_index_cache), ) -> WikiPageResponse: """Return the full content of a single wiki page identified by its title.""" - user_id = user.id if user else None + user_id = (user.id or None) if user else None wiki = await management.get_wiki(wiki_id, user_id=user_id) if wiki is None: raise HTTPException(404, f"Wiki not found: {wiki_id}") @@ -902,7 +910,7 @@ async def get_wiki_page( ) -> dict: """Get a single wiki page content. page_id can be section/page format.""" # Access control — verify caller can view this wiki - user_id = user.id if user else None + user_id = (user.id or None) if user else None wiki_record = await management.get_wiki(wiki_id, user_id=user_id) if wiki_record is None: raise HTTPException(404, f"Wiki not found: {wiki_id}") @@ -965,7 +973,7 @@ async def list_qa( status: QAStatus | None = Query(default=None), ) -> QAListResponse: """Paginated Q&A history for a wiki.""" - user_id = user.id if user else None + user_id = (user.id or None) if user else None wiki = await management.get_wiki(wiki_id, user_id=user_id) if not wiki: raise HTTPException(404, f"Wiki not found: {wiki_id}") @@ -982,7 +990,7 @@ async def qa_stats( management: WikiManagementService = Depends(get_wiki_management), ) -> QAStatsResponse: """QA statistics for a wiki.""" - user_id = user.id if user else None + user_id = (user.id or None) if user else None wiki = await management.get_wiki(wiki_id, user_id=user_id) if not wiki: raise HTTPException(404, f"Wiki not found: {wiki_id}") @@ -1064,7 +1072,7 @@ async def diff_wiki( # implementation detail of the repo, not shareable user data, so we # don't honor the wiki's shared-visibility flag for this endpoint # (unlike /search). Mirrors the /refresh guard. - user_id = user.id if user else None + user_id = (user.id or None) if user else None wiki = await management.get_wiki(wiki_id, user_id=user_id) if wiki is None: raise HTTPException(404, f"Wiki not found: {wiki_id}") @@ -1171,7 +1179,7 @@ async def incremental_refresh_wiki( Owner-only — same access model as ``/refresh`` since this surfaces internal content_hash + node_id state through the per-page events. """ - user_id = user.id if user else None + user_id = (user.id or None) if user else None wiki_record = await management.get_wiki(wiki_id, user_id=user_id) if wiki_record is None: raise HTTPException(404, f"Wiki not found: {wiki_id}") @@ -1231,7 +1239,7 @@ async def delete_wiki( management: WikiManagementService = Depends(get_wiki_management), service: WikiService = Depends(get_wiki_service), ) -> DeleteWikiResponse: - user_id = user.id if user else None + user_id = (user.id or None) if user else None settings = request.app.state.settings result = await management.delete_wiki(wiki_id, user_id=user_id, cache_dir=settings.cache_dir) @@ -1342,7 +1350,7 @@ async def export_wiki( 400, f"Invalid format '{format}'. Must be one of: {', '.join(sorted(_VALID_EXPORT_FORMATS))}" ) - user_id = user.id if user else None + user_id = (user.id or None) if user else None wiki_record = await management.get_wiki(wiki_id, user_id=user_id) if wiki_record is None: raise HTTPException(404, f"Wiki not found: {wiki_id}") @@ -1402,7 +1410,7 @@ async def import_wiki( import_service: ImportService = Depends(get_import_service), ) -> WikiSummary: """Import a wiki from a ``.wikiexport`` bundle produced by the wikis export.""" - user_id = user.id if user else None + user_id = (user.id or None) if user else None # Size guard — reject overly large uploads before reading if bundle.size is not None and bundle.size > _MAX_IMPORT_BYTES: @@ -1806,7 +1814,7 @@ async def project_codemap( Proxies to the codemap pipeline with ``project_id`` set and ``research_type=codemap``. """ - user_id = user.id if user else None + user_id = (user.id or None) if user else None # Verify the project exists and is accessible project = await svc.get_project(project_id, user_id=user_id or "") diff --git a/backend/tests/unit/test_get_wiki_is_owner.py b/backend/tests/unit/test_get_wiki_is_owner.py index 004619c..ee709d2 100644 --- a/backend/tests/unit/test_get_wiki_is_owner.py +++ b/backend/tests/unit/test_get_wiki_is_owner.py @@ -15,10 +15,11 @@ from unittest.mock import AsyncMock, MagicMock import pytest -from app.core.deep_research.research_engine import DeepResearchEngine from fastapi import FastAPI from httpx import ASGITransport, AsyncClient +from app.core.deep_research.research_engine import DeepResearchEngine + # --------------------------------------------------------------------------- # DeepResearchEngine._get_repo_context tests # --------------------------------------------------------------------------- @@ -309,7 +310,8 @@ async def test_get_wiki_override_clears_stale_error_when_refresh_is_live( staticmethod(lambda record, user_id: wiki_summary), ) async with AsyncClient( - transport=ASGITransport(app=app), base_url="http://test", + transport=ASGITransport(app=app), + base_url="http://test", ) as client: resp = await client.get("/api/v1/wikis/wiki-1") @@ -344,12 +346,15 @@ async def test_get_wiki_override_fires_for_running_incremental_refresh( mock_service.invocations = { # Stale terminal invocation comes first in iteration order. "inv-old-failed": _make_inflight_invocation( - status="failed", invocation_id="inv-old-failed", + status="failed", + invocation_id="inv-old-failed", ), # The live in-progress run is second — the loop must still # prefer it. "inv-r": _make_inflight_invocation( - status="running", progress=0.3, invocation_id="inv-r", + status="running", + progress=0.3, + invocation_id="inv-r", ), } @@ -362,7 +367,8 @@ async def test_get_wiki_override_fires_for_running_incremental_refresh( staticmethod(lambda record, user_id: wiki_summary), ) async with AsyncClient( - transport=ASGITransport(app=app), base_url="http://test", + transport=ASGITransport(app=app), + base_url="http://test", ) as client: resp = await client.get("/api/v1/wikis/wiki-1") @@ -399,7 +405,8 @@ async def test_get_wiki_terminal_invocation_does_not_override(app_with_mocks): staticmethod(lambda record, user_id: wiki_summary), ) async with AsyncClient( - transport=ASGITransport(app=app), base_url="http://test", + transport=ASGITransport(app=app), + base_url="http://test", ) as client: resp = await client.get("/api/v1/wikis/wiki-1") @@ -434,7 +441,8 @@ async def test_list_wikis_override_clears_stale_error_when_refresh_is_live( app, mock_management, mock_service = app_with_mocks wiki_summary = _make_serializable_wiki_summary( - status="failed", is_owner=True, + status="failed", + is_owner=True, error="Git clone failed (from prior attempt)", ) @@ -448,7 +456,8 @@ async def test_list_wikis_override_clears_stale_error_when_refresh_is_live( } async with AsyncClient( - transport=ASGITransport(app=app), base_url="http://test", + transport=ASGITransport(app=app), + base_url="http://test", ) as client: resp = await client.get("/api/v1/wikis") @@ -468,7 +477,8 @@ async def test_list_wikis_override_fires_for_running_incremental_refresh( app, mock_management, mock_service = app_with_mocks wiki_summary = _make_serializable_wiki_summary( - status="complete", is_owner=True, + status="complete", + is_owner=True, ) wiki_list = MagicMock() @@ -481,9 +491,83 @@ async def test_list_wikis_override_fires_for_running_incremental_refresh( } async with AsyncClient( - transport=ASGITransport(app=app), base_url="http://test", + transport=ASGITransport(app=app), + base_url="http://test", ) as client: resp = await client.get("/api/v1/wikis") assert resp.json()["wikis"][0]["status"] == "running" assert resp.json()["wikis"][0]["progress"] == 0.3 + + +# --------------------------------------------------------------------------- +# GET /api/v1/wikis — in-memory invocation ownership filter +# --------------------------------------------------------------------------- + + +def _make_foreign_invocation(wiki_id: str, owner_id: str, status: str = "generating") -> MagicMock: + inv = MagicMock() + inv.id = f"inv-{wiki_id}" + inv.wiki_id = wiki_id + inv.repo_url = f"https://github.com/example/{wiki_id}" + inv.branch = "main" + inv.status = status + inv.progress = 0.5 + inv.error = None + inv.pages_completed = 0 + inv.created_at = datetime(2024, 1, 3) + inv.owner_id = owner_id + return inv + + +@pytest.mark.asyncio +async def test_list_wikis_hides_other_users_in_progress_invocations(app_with_mocks): + """Invocations owned by other users must not appear in the caller's list.""" + app, mock_management, mock_service = app_with_mocks + + wiki_list = MagicMock() + wiki_list.wikis = [] + mock_management.list_wikis = AsyncMock(return_value=wiki_list) + mock_management.storage.list_artifacts = AsyncMock(return_value=[]) + + mock_service.invocations = { + "own": _make_foreign_invocation("wiki-own", owner_id="user-1"), + "other": _make_foreign_invocation("wiki-other", owner_id="user-2"), + "legacy": _make_foreign_invocation("wiki-legacy", owner_id=""), + } + + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + resp = await client.get("/api/v1/wikis") + + returned_ids = {w["wiki_id"] for w in resp.json()["wikis"]} + assert "wiki-own" in returned_ids, "caller's own invocation must be included" + assert "wiki-legacy" in returned_ids, "legacy unowned invocation must be included" + assert "wiki-other" not in returned_ids, "other user's invocation must be excluded" + + +@pytest.mark.asyncio +async def test_list_wikis_empty_string_user_id_does_not_leak_owned_invocations(app_with_mocks): + """user_id normalised to None when user.id is '' so the filter still runs.""" + app, mock_management, mock_service = app_with_mocks + + # Override the user dependency to return an empty-string id + from app.auth import get_current_user + + app.dependency_overrides[get_current_user] = lambda: MagicMock(id="") + + wiki_list = MagicMock() + wiki_list.wikis = [] + mock_management.list_wikis = AsyncMock(return_value=wiki_list) + mock_management.storage.list_artifacts = AsyncMock(return_value=[]) + + mock_service.invocations = { + "other": _make_foreign_invocation("wiki-other", owner_id="user-2"), + "legacy": _make_foreign_invocation("wiki-legacy", owner_id=""), + } + + async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client: + resp = await client.get("/api/v1/wikis") + + returned_ids = {w["wiki_id"] for w in resp.json()["wikis"]} + assert "wiki-other" not in returned_ids, "owned invocations must not leak when caller has empty user id" + assert "wiki-legacy" in returned_ids, "legacy unowned invocations still visible"