Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 27 additions & 19 deletions backend/app/api/routes.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Comment on lines +398 to 399
completed_wiki_ids = {w.wiki_id for w in result.wikis}

Expand All @@ -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
Expand All @@ -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
Comment on lines 444 to +446
Comment on lines +440 to +446
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:
Expand Down Expand Up @@ -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

Expand All @@ -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),
Expand Down Expand Up @@ -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}")
Expand Down Expand Up @@ -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}")
Expand Down Expand Up @@ -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}")
Expand Down Expand Up @@ -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}")
Expand Down Expand Up @@ -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}")
Expand Down Expand Up @@ -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}")
Expand All @@ -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}")
Expand Down Expand Up @@ -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}")
Expand Down Expand Up @@ -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}")
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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}")
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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 "")
Expand Down
104 changes: 94 additions & 10 deletions backend/tests/unit/test_get_wiki_is_owner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -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")

Expand Down Expand Up @@ -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",
),
}

Expand All @@ -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")

Expand Down Expand Up @@ -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")

Expand Down Expand Up @@ -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)",
)

Expand All @@ -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")

Expand All @@ -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()
Expand All @@ -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"
Comment on lines +542 to +545


@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=""),
}
Comment on lines +563 to +566

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"
Loading