From 80908ff807ad0afd61ad1f493dcc322f6c01d22a Mon Sep 17 00:00:00 2001 From: WenjingWang Date: Thu, 12 Mar 2026 15:49:20 +0100 Subject: [PATCH 1/9] fix(auth): replace user_id query param with JWT-based identity. Related to #151 --- src/paperbot/api/routes/gen_code.py | 30 ++- src/paperbot/api/routes/harvest.py | 24 +- src/paperbot/api/routes/memory.py | 16 +- src/paperbot/api/routes/repro_context.py | 17 +- src/paperbot/api/routes/research.py | 319 ++++++++++++----------- 5 files changed, 226 insertions(+), 180 deletions(-) diff --git a/src/paperbot/api/routes/gen_code.py b/src/paperbot/api/routes/gen_code.py index f0e398e7..e6cb767b 100644 --- a/src/paperbot/api/routes/gen_code.py +++ b/src/paperbot/api/routes/gen_code.py @@ -6,11 +6,13 @@ from pathlib import Path from typing import Optional -from fastapi import APIRouter, Request +from fastapi import APIRouter, Depends, Request +from fastapi.responses import StreamingResponse from pydantic import BaseModel from paperbot.application.collaboration.message_schema import new_run_id, new_trace_id from paperbot.core.abstractions import AgentRunContext +from paperbot.api.auth.dependencies import get_user_id from ..streaming import StreamEvent, sse_response @@ -181,7 +183,11 @@ async def gen_code_stream( @router.post("/gen-code") -async def generate_code(request: GenCodeRequest, http_request: Request): +async def generate_code( + request: GenCodeRequest, + http_request: Request, + user_id: str = Depends(get_user_id), +): """ Generate code from paper and stream progress. @@ -190,9 +196,19 @@ async def generate_code(request: GenCodeRequest, http_request: Request): event_log = getattr(http_request.app.state, "event_log", None) run_id = new_run_id() trace_id = new_trace_id() - return sse_response( - gen_code_stream(request, event_log=event_log, run_id=run_id, trace_id=trace_id), - workflow="gen_code", - run_id=run_id, - trace_id=trace_id, + # Respect authenticated user id where available; fall back to request.user_id + request.user_id = request.user_id or user_id + + return StreamingResponse( + wrap_generator( + gen_code_stream(request, event_log=event_log, run_id=run_id, trace_id=trace_id), + workflow="gen_code", + run_id=run_id, + trace_id=trace_id, + ), + media_type="text/event-stream", + headers={ + "Cache-Control": "no-cache", + "Connection": "keep-alive", + }, ) diff --git a/src/paperbot/api/routes/harvest.py b/src/paperbot/api/routes/harvest.py index 68a0e2e2..5ae94a7f 100644 --- a/src/paperbot/api/routes/harvest.py +++ b/src/paperbot/api/routes/harvest.py @@ -13,7 +13,8 @@ from typing import Any, Dict, List, Optional -from fastapi import APIRouter, HTTPException, Query, Request +from fastapi import APIRouter, Depends, HTTPException, Query, Request +from fastapi.responses import StreamingResponse from pydantic import BaseModel, Field from paperbot.api.streaming import StreamEvent, sse_response @@ -23,6 +24,7 @@ HarvestPipeline, HarvestProgress, ) +from paperbot.api.auth.dependencies import get_user_id from paperbot.infrastructure.stores.paper_store import PaperStore, paper_to_dict from paperbot.utils.logging_config import LogFiles, Logger, clear_trace_id, set_trace_id @@ -324,12 +326,9 @@ class LibraryResponse(BaseModel): offset: int -# TODO(auth): user_id is accepted from client without authentication. -# This is intentional for the MVP single-user setup. For multi-user production, -# user_id should come from an authenticated session or JWT token. @router.get("/papers/library", response_model=LibraryResponse) def get_user_library( - user_id: str = Query("default", description="User ID"), + user_id: str = Depends(get_user_id), track_id: Optional[int] = Query(None, description="Filter by track"), actions: Optional[str] = Query(None, description="Filter by actions (comma-separated)"), sort_by: str = Query("saved_at", description="Sort field"), @@ -392,13 +391,15 @@ def get_paper(paper_id: int): class SavePaperRequest(BaseModel): """Request to save paper to library.""" - # TODO(auth): user_id from client without auth - intentional for MVP single-user setup - user_id: str = Field("default", description="User ID") track_id: Optional[int] = Field(None, description="Associated track ID") @router.post("/papers/{paper_id}/save") -def save_paper_to_library(paper_id: int, request: SavePaperRequest): +def save_paper_to_library( + paper_id: int, + request: SavePaperRequest, + user_id: str = Depends(get_user_id), +): """ Save a paper to user's library. @@ -413,7 +414,7 @@ def save_paper_to_library(paper_id: int, request: SavePaperRequest): # Use research store to record feedback research_store = _get_research_store() feedback = research_store.record_paper_feedback( - user_id=request.user_id, + user_id=user_id, paper_id=str(paper_id), action="save", track_id=request.track_id, @@ -422,13 +423,10 @@ def save_paper_to_library(paper_id: int, request: SavePaperRequest): return {"success": True, "feedback": feedback} -# TODO(auth): user_id accepted from query string without authentication. -# Intentional for MVP single-user setup. For multi-user production, user_id -# should come from authenticated session/JWT, not query parameters. @router.delete("/papers/{paper_id}/save") def remove_paper_from_library( paper_id: int, - user_id: str = Query("default", description="User ID"), + user_id: str = Depends(get_user_id), ): """Remove a paper from user's library.""" store = _get_paper_store() diff --git a/src/paperbot/api/routes/memory.py b/src/paperbot/api/routes/memory.py index d81ac10f..65b5efd1 100644 --- a/src/paperbot/api/routes/memory.py +++ b/src/paperbot/api/routes/memory.py @@ -3,10 +3,11 @@ import hashlib from typing import Any, Dict, List, Optional -from fastapi import APIRouter, File, UploadFile, Query, HTTPException +from fastapi import APIRouter, File, UploadFile, Query, HTTPException, Depends from pydantic import BaseModel, Field from paperbot.infrastructure.stores.memory_store import SqlAlchemyMemoryStore +from paperbot.api.auth.dependencies import get_user_id from paperbot.memory import build_memory_context, extract_memories, parse_chat_log from paperbot.memory.schema import MemoryCandidate @@ -56,7 +57,7 @@ class IngestResponse(BaseModel): @router.post("/memory/ingest", response_model=IngestResponse) async def ingest_memory( file: UploadFile = File(...), - user_id: str = Query("default", description="Memory namespace; use one id per person/team."), + user_id: str = Depends(get_user_id), workspace_id: Optional[str] = Query(None, description="Optional workspace/project namespace."), platform: Optional[str] = Query(None, description="Hint: chatgpt/gemini/claude/..."), use_llm: bool = Query(False, description="Use configured LLM to extract memories (falls back on heuristics)."), @@ -122,7 +123,7 @@ class MemoryListResponse(BaseModel): @router.get("/memory/list", response_model=MemoryListResponse) def list_memories( - user_id: str = "default", + user_id: str = Depends(get_user_id), limit: int = 100, kind: Optional[str] = None, workspace_id: Optional[str] = None, @@ -174,9 +175,10 @@ class ContextResponse(BaseModel): @router.post("/memory/context", response_model=ContextResponse) -def memory_context(req: ContextRequest): - items = _get_store().search_memories( - user_id=req.user_id, +def memory_context(req: ContextRequest, user_id: str = Depends(get_user_id)): + effective_user_id = req.user_id or user_id + items = _store.search_memories( + user_id=effective_user_id, workspace_id=req.workspace_id, query=req.query, limit=req.limit, @@ -195,7 +197,7 @@ def memory_context(req: ContextRequest): ] ctx = build_memory_context(cands, max_items=req.limit) return ContextResponse( - user_id=req.user_id, + user_id=effective_user_id, query=req.query, context=ctx, items=[ diff --git a/src/paperbot/api/routes/repro_context.py b/src/paperbot/api/routes/repro_context.py index a061b266..7afbc172 100644 --- a/src/paperbot/api/routes/repro_context.py +++ b/src/paperbot/api/routes/repro_context.py @@ -16,10 +16,12 @@ from dataclasses import asdict as _asdict from typing import Literal, Optional -from fastapi import APIRouter, HTTPException +from fastapi import APIRouter, HTTPException, Depends +from fastapi.responses import StreamingResponse from pydantic import BaseModel -from paperbot.api.streaming import StreamEvent, sse_response +from paperbot.api.streaming import StreamEvent, wrap_generator +from paperbot.api.auth.dependencies import get_user_id from paperbot.application.services.p2c.models import ( GenerateContextRequest as P2CRequest, RawPaperData, @@ -330,11 +332,18 @@ async def _write_paper_scope_memories( @router.post("/generate") -async def generate_context_pack(request: GenerateContextPackRequest): +async def generate_context_pack( + request: GenerateContextPackRequest, + user_id: str = Depends(get_user_id), +): """Generate a P2C context pack for the given paper. Returns SSE stream.""" trace_id = set_trace_id() + # Prefer authenticated user id when available + effective_user_id = request.user_id or user_id + request.user_id = effective_user_id + Logger.info( - f"[M2] generate_request trace_id={trace_id} paper_id={request.paper_id} user_id={request.user_id}", + f"[M2] generate_request trace_id={trace_id} paper_id={request.paper_id} user_id={effective_user_id}", file=LogFiles.API, ) return sse_response(_generate_stream(request), workflow="p2c_generate") diff --git a/src/paperbot/api/routes/research.py b/src/paperbot/api/routes/research.py index 5475c848..5c14411d 100644 --- a/src/paperbot/api/routes/research.py +++ b/src/paperbot/api/routes/research.py @@ -7,7 +7,7 @@ from datetime import datetime, timedelta, timezone from typing import Any, Dict, List, Optional, Set, Tuple -from fastapi import APIRouter, BackgroundTasks, HTTPException, Query +from fastapi import APIRouter, BackgroundTasks, HTTPException, Query, Depends from fastapi.responses import Response from pydantic import BaseModel, Field from sqlalchemy.exc import IntegrityError @@ -36,6 +36,7 @@ from paperbot.infrastructure.stores.memory_store import SqlAlchemyMemoryStore from paperbot.infrastructure.stores.research_store import SqlAlchemyResearchStore from paperbot.infrastructure.stores.workflow_metric_store import WorkflowMetricStore +from paperbot.api.auth.dependencies import get_user_id from paperbot.memory.eval.collector import MemoryMetricCollector from paperbot.memory.extractor import extract_memories from paperbot.memory.schema import MemoryCandidate, NormalizedMessage @@ -370,7 +371,6 @@ def _collect_conference_deadline_terms(item: Dict[str, Any]) -> Set[str]: class TrackCreateRequest(BaseModel): - user_id: str = "default" name: str = Field(..., min_length=1, max_length=128) description: str = "" keywords: List[str] = [] @@ -438,9 +438,14 @@ def _serialize_track_context_response( @router.post("/research/tracks", response_model=TrackResponse) -def create_track(req: TrackCreateRequest, background_tasks: BackgroundTasks): - track = _get_research_store().create_track( - user_id=req.user_id, +def create_track( + req: TrackCreateRequest, + background_tasks: BackgroundTasks, + user_id: str = Depends(get_user_id), +): + effective_user_id = user_id + track = _research_store.create_track( + user_id=effective_user_id, name=req.name, description=req.description, keywords=req.keywords, @@ -464,7 +469,7 @@ class TrackListResponse(BaseModel): @router.get("/research/tracks", response_model=TrackListResponse) def list_tracks( - user_id: str = "default", + user_id: str = Depends(get_user_id), include_archived: bool = Query(False), limit: int = Query(100, ge=1, le=500), ): @@ -482,7 +487,7 @@ class DeadlineRadarResponse(BaseModel): @router.get("/research/deadlines/radar", response_model=DeadlineRadarResponse) def get_deadline_radar( - user_id: str = "default", + user_id: str = Depends(get_user_id), days: int = Query(180, ge=7, le=365), ccf_levels: str = Query("A,B,C"), field: Optional[str] = None, @@ -571,8 +576,8 @@ def get_deadline_radar( @router.get("/research/tracks/active", response_model=TrackResponse) -def get_active_track(user_id: str = "default"): - track = _get_research_store().get_active_track(user_id=user_id) +def get_active_track(user_id: str = Depends(get_user_id)): + track = _research_store.get_active_track(user_id=user_id) if not track: raise HTTPException(status_code=404, detail="No active track for user") return TrackResponse(track=track) @@ -594,7 +599,7 @@ def update_track( track_id: int, req: TrackUpdateRequest, background_tasks: BackgroundTasks, - user_id: str = "default", + user_id: str = Depends(get_user_id), ): update_data = req.model_dump(exclude_unset=True, exclude_none=True) @@ -619,8 +624,12 @@ def update_track( @router.post("/research/tracks/{track_id}/activate", response_model=TrackResponse) -def activate_track(track_id: int, background_tasks: BackgroundTasks, user_id: str = "default"): - track = _get_research_store().activate_track(user_id=user_id, track_id=track_id) +def activate_track( + track_id: int, + background_tasks: BackgroundTasks, + user_id: str = Depends(get_user_id), +): + track = _research_store.activate_track(user_id=user_id, track_id=track_id) if not track: raise HTTPException(status_code=404, detail="Track not found") _schedule_embedding_precompute(background_tasks, user_id=user_id, track_ids=[track_id]) @@ -628,7 +637,6 @@ def activate_track(track_id: int, background_tasks: BackgroundTasks, user_id: st class TaskCreateRequest(BaseModel): - user_id: str = "default" title: str = Field(..., min_length=1) status: str = "todo" priority: int = 0 @@ -642,9 +650,14 @@ class TaskResponse(BaseModel): @router.post("/research/tracks/{track_id}/tasks", response_model=TaskResponse) -def add_task(track_id: int, req: TaskCreateRequest, background_tasks: BackgroundTasks): - task = _get_research_store().add_task( - user_id=req.user_id, +def add_task( + track_id: int, + req: TaskCreateRequest, + background_tasks: BackgroundTasks, + user_id: str = Depends(get_user_id), +): + task = _research_store.add_task( + user_id=user_id, track_id=track_id, title=req.title, status=req.status, @@ -655,7 +668,7 @@ def add_task(track_id: int, req: TaskCreateRequest, background_tasks: Background ) if not task: raise HTTPException(status_code=404, detail="Track not found") - _schedule_embedding_precompute(background_tasks, user_id=req.user_id, track_ids=[track_id]) + _schedule_embedding_precompute(background_tasks, user_id=user_id, track_ids=[track_id]) return TaskResponse(task=task) @@ -668,7 +681,7 @@ class TaskListResponse(BaseModel): @router.get("/research/tracks/{track_id}/tasks", response_model=TaskListResponse) def list_tasks( track_id: int, - user_id: str = "default", + user_id: str = Depends(get_user_id), status: Optional[str] = None, limit: int = Query(100, ge=1, le=500), ): @@ -679,7 +692,6 @@ def list_tasks( class MemoryItemCreateRequest(BaseModel): - user_id: str = "default" scope_type: str = "track" # global/track scope_id: Optional[str] = None kind: str = Field("note", min_length=1, max_length=32) @@ -705,9 +717,13 @@ def _resolve_track_scope_id( @router.post("/research/memory/items", response_model=MemoryItemResponse) -def create_memory_item(req: MemoryItemCreateRequest, background_tasks: BackgroundTasks): +def create_memory_item( + req: MemoryItemCreateRequest, + background_tasks: BackgroundTasks, + user_id: str = Depends(get_user_id), +): scope_type = (req.scope_type or "global").strip() or "global" - scope_id = _resolve_track_scope_id(req.user_id, scope_type, req.scope_id) + scope_id = _resolve_track_scope_id(user_id, scope_type, req.scope_id) if scope_type == "track" and not scope_id: raise HTTPException( status_code=400, @@ -724,14 +740,14 @@ def create_memory_item(req: MemoryItemCreateRequest, background_tasks: Backgroun scope_id=scope_id, status=req.status, ) - created, _, rows = _get_memory_store().add_memories(user_id=req.user_id, memories=[cand]) + created, _, rows = _memory_store.add_memories(user_id=user_id, memories=[cand]) if created <= 0 or not rows: raise HTTPException( status_code=409, detail="Duplicate memory item (same scope/kind/content)" ) if scope_type == "track": _schedule_embedding_precompute( - background_tasks, user_id=req.user_id, track_ids=[int(scope_id or 0)] + background_tasks, user_id=user_id, track_ids=[int(scope_id or 0)] ) return MemoryItemResponse(item=SqlAlchemyMemoryStore._row_to_dict(rows[0])) @@ -743,7 +759,7 @@ class MemoryItemListResponse(BaseModel): @router.get("/research/memory/items", response_model=MemoryItemListResponse) def list_memory_items( - user_id: str = "default", + user_id: str = Depends(get_user_id), scope_type: Optional[str] = None, scope_id: Optional[str] = None, kind: Optional[str] = None, @@ -765,7 +781,7 @@ def list_memory_items( @router.get("/research/memory/inbox", response_model=MemoryItemListResponse) def list_memory_inbox( - user_id: str = "default", + user_id: str = Depends(get_user_id), track_id: Optional[int] = None, limit: int = Query(100, ge=1, le=500), ): @@ -781,7 +797,6 @@ def list_memory_inbox( class MemorySuggestRequest(BaseModel): - user_id: str = "default" text: str = Field(..., min_length=1) scope_type: str = "track" scope_id: Optional[str] = None @@ -798,9 +813,13 @@ class MemorySuggestResponse(BaseModel): @router.post("/research/memory/suggest", response_model=MemorySuggestResponse) -def suggest_memories(req: MemorySuggestRequest, background_tasks: BackgroundTasks): +def suggest_memories( + req: MemorySuggestRequest, + background_tasks: BackgroundTasks, + user_id: str = Depends(get_user_id), +): scope_type = (req.scope_type or "global").strip() or "global" - scope_id = _resolve_track_scope_id(req.user_id, scope_type, req.scope_id) + scope_id = _resolve_track_scope_id(user_id, scope_type, req.scope_id) if scope_type == "track" and not scope_id: raise HTTPException( status_code=400, @@ -824,13 +843,13 @@ def suggest_memories(req: MemorySuggestRequest, background_tasks: BackgroundTask ) for m in extracted ] - created, skipped, rows = _get_memory_store().add_memories(user_id=req.user_id, memories=pending) + created, skipped, rows = _memory_store.add_memories(user_id=user_id, memories=pending) if scope_type == "track": _schedule_embedding_precompute( - background_tasks, user_id=req.user_id, track_ids=[int(scope_id or 0)] + background_tasks, user_id=user_id, track_ids=[int(scope_id or 0)] ) return MemorySuggestResponse( - user_id=req.user_id, + user_id=user_id, created=created, skipped=skipped, candidates=[SqlAlchemyMemoryStore._row_to_dict(r) for r in rows], @@ -838,7 +857,6 @@ def suggest_memories(req: MemorySuggestRequest, background_tasks: BackgroundTask class MemoryModerateRequest(BaseModel): - user_id: str = "default" status: str = Field(..., min_length=1) # approved/rejected/pending/superseded content: Optional[str] = None kind: Optional[str] = None @@ -848,9 +866,13 @@ class MemoryModerateRequest(BaseModel): @router.post("/research/memory/items/{item_id}/moderate", response_model=MemoryItemResponse) -def moderate_memory_item(item_id: int, req: MemoryModerateRequest): - updated = _get_memory_store().update_item( - user_id=req.user_id, +def moderate_memory_item( + item_id: int, + req: MemoryModerateRequest, + user_id: str = Depends(get_user_id), +): + updated = _memory_store.update_item( + user_id=user_id, item_id=item_id, status=req.status, content=req.content, @@ -866,7 +888,6 @@ def moderate_memory_item(item_id: int, req: MemoryModerateRequest): class BulkModerateRequest(BaseModel): - user_id: str = "default" item_ids: List[int] = Field(default_factory=list) status: str = Field(..., min_length=1) @@ -877,9 +898,13 @@ class BulkModerateResponse(BaseModel): @router.post("/research/memory/bulk_moderate", response_model=BulkModerateResponse) -def bulk_moderate(req: BulkModerateRequest, background_tasks: BackgroundTasks): +def bulk_moderate( + req: BulkModerateRequest, + background_tasks: BackgroundTasks, + user_id: str = Depends(get_user_id), +): result = _build_track_memory_service().bulk_moderate( - user_id=req.user_id, + user_id=user_id, item_ids=req.item_ids, status=req.status, ) @@ -887,7 +912,7 @@ def bulk_moderate(req: BulkModerateRequest, background_tasks: BackgroundTasks): updated = result.updated_items _schedule_embedding_precompute( background_tasks, - user_id=req.user_id, + user_id=user_id, track_ids=result.affected_track_ids, ) @@ -904,18 +929,17 @@ def bulk_moderate(req: BulkModerateRequest, background_tasks: BackgroundTasks): collector.record_false_positive_rate( false_positive_count=high_confidence_rejected, total_approved_count=len(items_before), - evaluator_id=f"user:{req.user_id}", + evaluator_id=f"user:{user_id}", detail={ "item_ids": req.item_ids, "action": "bulk_moderate_reject", }, ) - return BulkModerateResponse(user_id=req.user_id, updated=updated) + return BulkModerateResponse(user_id=user_id, updated=updated) class BulkMoveRequest(BaseModel): - user_id: str = "default" item_ids: List[int] = Field(default_factory=list) scope_type: str = Field(..., min_length=1) scope_id: Optional[str] = None @@ -927,10 +951,14 @@ class BulkMoveResponse(BaseModel): @router.post("/research/memory/bulk_move", response_model=BulkMoveResponse) -def bulk_move(req: BulkMoveRequest, background_tasks: BackgroundTasks): +def bulk_move( + req: BulkMoveRequest, + background_tasks: BackgroundTasks, + user_id: str = Depends(get_user_id), +): try: result = _build_track_memory_service().bulk_move( - user_id=req.user_id, + user_id=user_id, item_ids=req.item_ids, scope_type=req.scope_type, scope_id=req.scope_id, @@ -940,16 +968,15 @@ def bulk_move(req: BulkMoveRequest, background_tasks: BackgroundTasks): _schedule_embedding_precompute( background_tasks, - user_id=req.user_id, + user_id=user_id, track_ids=result.affected_track_ids, ) - return BulkMoveResponse(user_id=req.user_id, updated=result.updated_items) + return BulkMoveResponse(user_id=user_id, updated=result.updated_items) class MemoryFeedbackRequest(BaseModel): """Request to record feedback on retrieved memories.""" - user_id: str = "default" memory_ids: List[int] = Field(..., min_length=1, description="IDs of memories being rated") helpful_ids: List[int] = Field( default_factory=list, description="IDs of memories that were helpful" @@ -970,7 +997,10 @@ class MemoryFeedbackResponse(BaseModel): @router.post("/research/memory/feedback", response_model=MemoryFeedbackResponse) -def record_memory_feedback(req: MemoryFeedbackRequest): +def record_memory_feedback( + req: MemoryFeedbackRequest, + user_id: str = Depends(get_user_id), +): """ Record user feedback on retrieved memories. @@ -995,7 +1025,7 @@ def record_memory_feedback(req: MemoryFeedbackRequest): collector.record_retrieval_hit_rate( hits=hits, expected=total, - evaluator_id=f"user:{req.user_id}", + evaluator_id=f"user:{user_id}", detail={ "memory_ids": req.memory_ids, "helpful_ids": list(helpful_set), @@ -1009,7 +1039,7 @@ def record_memory_feedback(req: MemoryFeedbackRequest): hit_rate = 0.0 return MemoryFeedbackResponse( - user_id=req.user_id, + user_id=user_id, total_rated=total, helpful_count=hits, not_helpful_count=len(not_helpful_set), @@ -1027,7 +1057,7 @@ class ClearTrackMemoryResponse(BaseModel): def clear_track_memory( track_id: int, background_tasks: BackgroundTasks, - user_id: str = "default", + user_id: str = Depends(get_user_id), confirm: bool = Query(False), ): if not confirm: @@ -1062,7 +1092,6 @@ def clear_track_memory( class PrecomputeEmbeddingsRequest(BaseModel): - user_id: str = "default" track_ids: Optional[List[int]] = None @@ -1072,11 +1101,14 @@ class PrecomputeEmbeddingsResponse(BaseModel): @router.post("/research/embeddings/precompute", response_model=PrecomputeEmbeddingsResponse) -def precompute_embeddings(req: PrecomputeEmbeddingsRequest): - result = _get_track_router().precompute_track_embeddings( - user_id=req.user_id, track_ids=req.track_ids or None +def precompute_embeddings( + req: PrecomputeEmbeddingsRequest, + user_id: str = Depends(get_user_id), +): + result = _track_router.precompute_track_embeddings( + user_id=user_id, track_ids=req.track_ids or None ) - return PrecomputeEmbeddingsResponse(user_id=req.user_id, result=result) + return PrecomputeEmbeddingsResponse(user_id=user_id, result=result) class EvalSummaryResponse(BaseModel): @@ -1087,7 +1119,7 @@ class EvalSummaryResponse(BaseModel): @router.get("/research/evals/summary", response_model=EvalSummaryResponse) def eval_summary( - user_id: str = "default", + user_id: str = Depends(get_user_id), track_id: Optional[int] = None, days: int = Query(30, ge=1, le=365), ): @@ -1096,7 +1128,6 @@ def eval_summary( class PaperFeedbackRequest(BaseModel): - user_id: str = "default" track_id: Optional[int] = None paper_id: str = Field(..., min_length=1) action: str = Field(..., min_length=1) # like/unlike/dislike/undislike/skip/save/unsave/cite @@ -1122,7 +1153,11 @@ class PaperFeedbackResponse(BaseModel): @router.post("/research/papers/feedback", response_model=PaperFeedbackResponse) -def add_paper_feedback(req: PaperFeedbackRequest, background_tasks: BackgroundTasks): +def add_paper_feedback( + req: PaperFeedbackRequest, + background_tasks: BackgroundTasks, + user_id: str = Depends(get_user_id), +): set_trace_id() # Initialize trace_id for this request Logger.info(f"Received paper feedback request, action={req.action}", file=LogFiles.HARVEST) research_store = _get_research_store() @@ -1131,7 +1166,7 @@ def add_paper_feedback(req: PaperFeedbackRequest, background_tasks: BackgroundTa active_track: Optional[Dict[str, Any]] = None if track_id is None: Logger.info("No track specified, getting active track", file=LogFiles.HARVEST) - active_track = research_store.get_active_track(user_id=req.user_id) + active_track = research_store.get_active_track(user_id=user_id) if not active_track: Logger.error("No active track found", file=LogFiles.HARVEST) raise HTTPException(status_code=400, detail="track_id missing and no active track") @@ -1195,8 +1230,8 @@ def add_paper_feedback(req: PaperFeedbackRequest, background_tasks: BackgroundTa Logger.warning(f"Failed to save paper to library: {e}", file=LogFiles.HARVEST) Logger.info("Recording paper feedback to research store", file=LogFiles.HARVEST) - fb = research_store.add_paper_feedback( - user_id=req.user_id, + fb = _research_store.add_paper_feedback( + user_id=user_id, track_id=track_id, paper_id=req.paper_id, # Always use external ID for consistency action=req.action, @@ -1232,7 +1267,7 @@ class PaperFeedbackListResponse(BaseModel): @router.get("/research/tracks/{track_id}/papers/feedback", response_model=PaperFeedbackListResponse) def list_paper_feedback( track_id: int, - user_id: str = "default", + user_id: str = Depends(get_user_id), action: Optional[str] = None, limit: int = Query(200, ge=1, le=1000), ): @@ -1243,7 +1278,6 @@ def list_paper_feedback( class PaperReadingStatusRequest(BaseModel): - user_id: str = "default" status: str = Field(..., min_length=1) # unread/reading/read/archived mark_saved: Optional[bool] = None metadata: Dict[str, Any] = {} @@ -1259,7 +1293,6 @@ class SavedPapersResponse(BaseModel): class DiscoverySeedRequest(BaseModel): - user_id: str = "default" track_id: Optional[int] = None seed_type: str = Field(..., pattern="^(doi|arxiv|openalex|semantic_scholar|author)$") seed_id: str = Field(..., min_length=1) @@ -1282,28 +1315,24 @@ class DiscoverySeedResponse(BaseModel): class PaperCollectionCreateRequest(BaseModel): - user_id: str = "default" name: str = Field(..., min_length=1, max_length=128) description: str = "" track_id: Optional[int] = Field(default=None, ge=1) class PaperCollectionUpdateRequest(BaseModel): - user_id: str = "default" name: Optional[str] = Field(default=None, min_length=1, max_length=128) description: Optional[str] = None archived: Optional[bool] = None class PaperCollectionItemUpsertRequest(BaseModel): - user_id: str = "default" paper_id: str = Field(..., min_length=1) note: Optional[str] = "" tags: Optional[List[str]] = [] class PaperCollectionItemPatchRequest(BaseModel): - user_id: str = "default" note: Optional[str] = "" tags: Optional[List[str]] = [] @@ -1342,7 +1371,6 @@ class AnchorDiscoverResponse(BaseModel): class AnchorActionRequest(BaseModel): - user_id: str = "default" action: str = Field(..., pattern="^(follow|ignore)$") @@ -1366,9 +1394,13 @@ class PaperRepoListResponse(BaseModel): @router.post("/research/papers/{paper_id}/status", response_model=PaperReadingStatusResponse) -def update_paper_status(paper_id: str, req: PaperReadingStatusRequest): - status = _get_research_store().set_paper_reading_status( - user_id=req.user_id, +def update_paper_status( + paper_id: str, + req: PaperReadingStatusRequest, + user_id: str = Depends(get_user_id), +): + status = _research_store.set_paper_reading_status( + user_id=user_id, paper_id=paper_id, status=req.status, metadata=req.metadata, @@ -1381,7 +1413,7 @@ def update_paper_status(paper_id: str, req: PaperReadingStatusRequest): @router.get("/research/papers/saved", response_model=SavedPapersResponse) def list_saved_papers( - user_id: str = "default", + user_id: str = Depends(get_user_id), track_id: Optional[int] = None, collection_id: Optional[int] = None, sort_by: str = Query("saved_at"), @@ -1398,7 +1430,7 @@ def list_saved_papers( @router.post("/research/discovery/seed", response_model=DiscoverySeedResponse) -async def discover_from_seed(req: DiscoverySeedRequest): +async def discover_from_seed(req: DiscoverySeedRequest, user_id: str = Depends(get_user_id)): if req.year_from and req.year_to and req.year_from > req.year_to: raise HTTPException(status_code=400, detail="year_from must be <= year_to") @@ -1577,7 +1609,7 @@ async def discover_from_seed(req: DiscoverySeedRequest): year_to=req.year_to, ) feedback_profile = ( - _build_feedback_profile(user_id=req.user_id, track_id=req.track_id) + _build_feedback_profile(user_id=user_id, track_id=req.track_id) if req.personalized else {} ) @@ -1649,10 +1681,10 @@ async def discover_from_seed(req: DiscoverySeedRequest): @router.post("/research/collections", response_model=PaperCollectionResponse) -def create_collection(req: PaperCollectionCreateRequest): +def create_collection(req: PaperCollectionCreateRequest, user_id: str = Depends(get_user_id)): try: - collection = _get_research_store().create_collection( - user_id=req.user_id, + collection = _research_store.create_collection( + user_id=user_id, name=req.name, description=req.description, track_id=req.track_id, @@ -1664,7 +1696,7 @@ def create_collection(req: PaperCollectionCreateRequest): @router.get("/research/collections", response_model=PaperCollectionListResponse) def list_collections( - user_id: str = "default", + user_id: str = Depends(get_user_id), include_archived: bool = Query(False), track_id: Optional[int] = Query(default=None), limit: int = Query(200, ge=1, le=1000), @@ -1679,10 +1711,10 @@ def list_collections( @router.patch("/research/collections/{collection_id}", response_model=PaperCollectionResponse) -def update_collection(collection_id: int, req: PaperCollectionUpdateRequest): +def update_collection(collection_id: int, req: PaperCollectionUpdateRequest, user_id: str = Depends(get_user_id)): try: - collection = _get_research_store().update_collection( - user_id=req.user_id, + collection = _research_store.update_collection( + user_id=user_id, collection_id=collection_id, name=req.name, description=req.description, @@ -1701,7 +1733,7 @@ def update_collection(collection_id: int, req: PaperCollectionUpdateRequest): ) def list_collection_items( collection_id: int, - user_id: str = "default", + user_id: str = Depends(get_user_id), limit: int = Query(500, ge=1, le=5000), ): items = _get_research_store().list_collection_items( @@ -1716,9 +1748,9 @@ def list_collection_items( "/research/collections/{collection_id}/items", response_model=PaperCollectionItemsResponse, ) -def upsert_collection_item(collection_id: int, req: PaperCollectionItemUpsertRequest): - item = _get_research_store().upsert_collection_item( - user_id=req.user_id, +def upsert_collection_item(collection_id: int, req: PaperCollectionItemUpsertRequest, user_id: str = Depends(get_user_id)): + item = _research_store.upsert_collection_item( + user_id=user_id, collection_id=collection_id, paper_id=req.paper_id, note=req.note, @@ -1726,11 +1758,9 @@ def upsert_collection_item(collection_id: int, req: PaperCollectionItemUpsertReq ) if item is None: raise HTTPException(status_code=404, detail="Collection or paper not found") - items = _get_research_store().list_collection_items( - user_id=req.user_id, collection_id=collection_id - ) + items = _research_store.list_collection_items(user_id=user_id, collection_id=collection_id) return PaperCollectionItemsResponse( - user_id=req.user_id, collection_id=collection_id, items=items + user_id=user_id, collection_id=collection_id, items=items ) @@ -1738,9 +1768,9 @@ def upsert_collection_item(collection_id: int, req: PaperCollectionItemUpsertReq "/research/collections/{collection_id}/items/{paper_id}", response_model=PaperCollectionItemsResponse, ) -def patch_collection_item(collection_id: int, paper_id: str, req: PaperCollectionItemPatchRequest): - item = _get_research_store().upsert_collection_item( - user_id=req.user_id, +def patch_collection_item(collection_id: int, paper_id: str, req: PaperCollectionItemPatchRequest, user_id: str = Depends(get_user_id)): + item = _research_store.upsert_collection_item( + user_id=user_id, collection_id=collection_id, paper_id=paper_id, note=req.note, @@ -1748,17 +1778,15 @@ def patch_collection_item(collection_id: int, paper_id: str, req: PaperCollectio ) if item is None: raise HTTPException(status_code=404, detail="Collection or paper not found") - items = _get_research_store().list_collection_items( - user_id=req.user_id, collection_id=collection_id - ) + items = _research_store.list_collection_items(user_id=user_id, collection_id=collection_id) return PaperCollectionItemsResponse( - user_id=req.user_id, collection_id=collection_id, items=items + user_id=user_id, collection_id=collection_id, items=items ) @router.delete("/research/collections/{collection_id}/items/{paper_id}") -def delete_collection_item(collection_id: int, paper_id: str, user_id: str = "default"): - ok = _get_research_store().remove_collection_item( +def delete_collection_item(collection_id: int, paper_id: str, user_id: str = Depends(get_user_id)): + ok = _research_store.remove_collection_item( user_id=user_id, collection_id=collection_id, paper_id=paper_id, @@ -1771,7 +1799,7 @@ def delete_collection_item(collection_id: int, paper_id: str, user_id: str = "de @router.get("/research/tracks/{track_id}/feed", response_model=TrackFeedResponse) def get_track_feed( track_id: int, - user_id: str = "default", + user_id: str = Depends(get_user_id), limit: int = Query(20, ge=1, le=100), offset: int = Query(0, ge=0), ): @@ -1797,7 +1825,7 @@ def get_track_feed( @router.get("/research/papers/export") def export_papers( - user_id: str = "default", + user_id: str = Depends(get_user_id), track_id: Optional[int] = None, format: str = Query("bibtex", pattern="^(bibtex|ris|markdown|csl_json)$"), ): @@ -1848,7 +1876,6 @@ def export_papers( class BibtexImportRequest(BaseModel): - user_id: str = "default" content: str = Field(..., min_length=1) track_id: Optional[int] = Field(default=None, ge=1) track_name: Optional[str] = Field(default=None, min_length=1, max_length=128) @@ -1868,13 +1895,13 @@ class BibtexImportResponse(BaseModel): @router.post("/research/papers/import/bibtex", response_model=BibtexImportResponse) -def import_bibtex(req: BibtexImportRequest): +def import_bibtex(req: BibtexImportRequest, user_id: str = Depends(get_user_id)): entries = _parse_bibtex_entries(req.content) if not entries: raise HTTPException(status_code=400, detail="No valid BibTeX entries found") track = _resolve_or_create_import_track( - user_id=req.user_id, + user_id=user_id, track_id=req.track_id, track_name=req.track_name, default_track_name="BibTeX Imports", @@ -1882,8 +1909,8 @@ def import_bibtex(req: BibtexImportRequest): track_pk = int(track["id"]) paper_store = _get_paper_store() - existing_saved_ids = _get_research_store().list_paper_feedback_ids( - user_id=req.user_id, + existing_saved_ids = _research_store.list_paper_feedback_ids( + user_id=user_id, track_id=track_pk, action="save", limit=5000, @@ -1923,8 +1950,8 @@ def import_bibtex(req: BibtexImportRequest): "citation_key": str(entry.get("key") or ""), "entry_type": str(entry.get("entry_type") or ""), } - _get_research_store().add_paper_feedback( - user_id=req.user_id, + _research_store.add_paper_feedback( + user_id=user_id, track_id=track_pk, paper_id=paper_ref, action="save", @@ -1937,7 +1964,7 @@ def import_bibtex(req: BibtexImportRequest): errors.append(f"entry {index}: {exc}") return BibtexImportResponse( - user_id=req.user_id, + user_id=user_id, track_id=track_pk, track_name=str(track.get("name") or ""), parsed=len(entries), @@ -1950,7 +1977,6 @@ def import_bibtex(req: BibtexImportRequest): class ZoteroSyncRequest(BaseModel): - user_id: str = "default" track_id: Optional[int] = Field(default=None, ge=1) track_name: Optional[str] = Field(default=None, min_length=1, max_length=128) library_type: str = Field(default="user", pattern="^(user|group)$") @@ -1972,11 +1998,11 @@ class ZoteroPullResponse(BaseModel): @router.post("/research/integrations/zotero/pull", response_model=ZoteroPullResponse) -def pull_from_zotero(req: ZoteroSyncRequest): +def pull_from_zotero(req: ZoteroSyncRequest, user_id: str = Depends(get_user_id)): from paperbot.infrastructure.connectors.zotero_connector import ZoteroConnector track = _resolve_or_create_import_track( - user_id=req.user_id, + user_id=user_id, track_id=req.track_id, track_name=req.track_name, default_track_name="Zotero Imports", @@ -1995,8 +2021,8 @@ def pull_from_zotero(req: ZoteroSyncRequest): raise HTTPException(status_code=502, detail=f"Failed to pull from Zotero: {exc}") from exc paper_store = _get_paper_store() - existing_saved_ids = _get_research_store().list_paper_feedback_ids( - user_id=req.user_id, + existing_saved_ids = _research_store.list_paper_feedback_ids( + user_id=user_id, track_id=track_pk, action="save", limit=5000, @@ -2036,8 +2062,8 @@ def pull_from_zotero(req: ZoteroSyncRequest): "zotero_library_type": req.library_type, "zotero_library_id": req.library_id, } - _get_research_store().add_paper_feedback( - user_id=req.user_id, + _research_store.add_paper_feedback( + user_id=user_id, track_id=track_pk, paper_id=paper_ref, action="save", @@ -2050,7 +2076,7 @@ def pull_from_zotero(req: ZoteroSyncRequest): errors.append(f"item {index}: {exc}") return ZoteroPullResponse( - user_id=req.user_id, + user_id=user_id, track_id=track_pk, track_name=str(track.get("name") or ""), total_remote=len(remote_items), @@ -2079,17 +2105,17 @@ class ZoteroPushResponse(BaseModel): @router.post("/research/integrations/zotero/push", response_model=ZoteroPushResponse) -def push_to_zotero(req: ZoteroPushRequest): +def push_to_zotero(req: ZoteroPushRequest, user_id: str = Depends(get_user_id)): from paperbot.infrastructure.connectors.zotero_connector import ZoteroConnector if req.track_id is not None: - track = _get_research_store().get_track(user_id=req.user_id, track_id=req.track_id) + track = _research_store.get_track(user_id=user_id, track_id=req.track_id) if not track: raise HTTPException(status_code=404, detail="Track not found") connector = ZoteroConnector() - local_items = _get_research_store().list_saved_papers( - user_id=req.user_id, + local_items = _research_store.list_saved_papers( + user_id=user_id, track_id=req.track_id, sort_by="saved_at", limit=req.max_items, @@ -2146,7 +2172,7 @@ def push_to_zotero(req: ZoteroPushRequest): errors.append(f"batch {start // 50 + 1}: {exc}") return ZoteroPushResponse( - user_id=req.user_id, + user_id=user_id, track_id=req.track_id, local_saved=len(local_papers), remote_items=len(remote_items), @@ -2159,10 +2185,9 @@ def push_to_zotero(req: ZoteroPushRequest): @router.get("/research/tracks/{track_id}/anchors/discover", response_model=AnchorDiscoverResponse) -# TODO: IDOR — replace user_id query param with authenticated session user (PR #112 review). def discover_track_anchors( track_id: int, - user_id: str = "default", + user_id: str = Depends(get_user_id), limit: int = Query(20, ge=1, le=100), window_days: int = Query(365, ge=30, le=3650), personalized: bool = Query(True), @@ -2199,18 +2224,17 @@ def discover_track_anchors( "/research/tracks/{track_id}/anchors/{author_id}/action", response_model=AnchorActionResponse, ) -# TODO: IDOR — replace req.user_id with authenticated session user (PR #112 review). -def set_anchor_action(track_id: int, author_id: int, req: AnchorActionRequest): +def set_anchor_action(track_id: int, author_id: int, req: AnchorActionRequest, user_id: str = Depends(get_user_id)): _ensure_anchor_feature_enabled() - track = _get_research_store().get_track(user_id=req.user_id, track_id=track_id) + track = _research_store.get_track(user_id=user_id, track_id=track_id) if not track: raise HTTPException(status_code=404, detail="Track not found") started = time.perf_counter() try: payload = _get_anchor_service().set_user_anchor_action( - user_id=req.user_id, + user_id=user_id, track_id=track_id, author_id=author_id, action=req.action, @@ -2229,14 +2253,14 @@ def set_anchor_action(track_id: int, author_id: int, req: AnchorActionRequest): claim_count=0, evidence_count=0, elapsed_ms=elapsed_ms, - detail={"author_id": author_id, "user_id": req.user_id}, + detail={"author_id": author_id, "user_id": user_id}, ) return AnchorActionResponse(action=payload) @router.get("/research/tracks/{track_id}/anchors/actions", response_model=AnchorActionListResponse) -def list_anchor_actions(track_id: int, user_id: str = "default"): +def list_anchor_actions(track_id: int, user_id: str = Depends(get_user_id)): _ensure_anchor_feature_enabled() track = _get_research_store().get_track(user_id=user_id, track_id=track_id) @@ -2248,8 +2272,8 @@ def list_anchor_actions(track_id: int, user_id: str = "default"): @router.get("/research/papers/{paper_id}", response_model=PaperDetailResponse) -def get_paper_detail(paper_id: str, user_id: str = "default"): - detail = _get_research_store().get_paper_detail(paper_id=paper_id, user_id=user_id) +def get_paper_detail(paper_id: str, user_id: str = Depends(get_user_id)): + detail = _research_store.get_paper_detail(paper_id=paper_id, user_id=user_id) if not detail: raise HTTPException(status_code=404, detail="Paper not found in registry") return PaperDetailResponse(detail=detail) @@ -2264,7 +2288,6 @@ def get_paper_repos(paper_id: str): class RouterSuggestRequest(BaseModel): - user_id: str = "default" query: str = Field(..., min_length=1) @@ -2273,19 +2296,19 @@ class RouterSuggestResponse(BaseModel): @router.post("/research/router/suggest", response_model=RouterSuggestResponse) -def suggest_track(req: RouterSuggestRequest): - active = _get_research_store().get_active_track(user_id=req.user_id) +def suggest_track(req: RouterSuggestRequest, user_id: str = Depends(get_user_id)): + active = _research_store.get_active_track(user_id=user_id) if not active: return RouterSuggestResponse(suggestion=None) grounded_query = _get_workflow_query_grounder().ground_query( - user_id=req.user_id, query=req.query + user_id=user_id, query=req.query ) routing_query = build_grounded_routing_query( original_query=req.query, grounded_query=grounded_query, ) suggestion = _get_track_router().suggest_track( - user_id=req.user_id, + user_id=user_id, query=routing_query or req.query, active_track_id=int(active["id"]), ) @@ -2298,7 +2321,6 @@ def suggest_track(req: RouterSuggestRequest): class ContextRequest(BaseModel): - user_id: str = "default" query: str = Field(..., min_length=1) track_id: Optional[int] = None activate_track_id: Optional[int] = None # confirm switch: activates then uses it @@ -2378,7 +2400,7 @@ def get_evidence_coverage( @router.post("/research/context", response_model=ContextResponse) -async def build_context(req: ContextRequest): +async def build_context(req: ContextRequest, user_id: str = Depends(get_user_id)): set_trace_id() # Initialize trace_id for this request Logger.info("Received build context request", file=LogFiles.HARVEST) @@ -2390,8 +2412,8 @@ async def build_context(req: ContextRequest): if req.activate_track_id is not None: Logger.info("Activating research track", file=LogFiles.HARVEST) - activated = _get_research_store().activate_track( - user_id=req.user_id, track_id=req.activate_track_id + activated = _research_store.activate_track( + user_id=user_id, track_id=req.activate_track_id ) if not activated: Logger.error("Research track not found", file=LogFiles.HARVEST) @@ -2436,7 +2458,7 @@ async def build_context(req: ContextRequest): try: Logger.info("Building context pack with paper recommendations", file=LogFiles.HARVEST) pack = await engine.build_context_pack( - user_id=req.user_id, + user_id=user_id, query=req.query, track_id=req.track_id, include_cross_track=req.include_cross_track, @@ -3960,8 +3982,8 @@ class StructuredCardResponse(BaseModel): @router.get("/research/papers/{paper_id}/card", response_model=StructuredCardResponse) -def get_structured_card(paper_id: str, user_id: str = "default"): - detail = _get_research_store().get_paper_detail(paper_id=paper_id, user_id=user_id) +def get_structured_card(paper_id: str, user_id: str = Depends(get_user_id)): + detail = _research_store.get_paper_detail(paper_id=paper_id, user_id=user_id) if not detail: raise HTTPException(status_code=404, detail="Paper not found") @@ -4012,7 +4034,6 @@ def get_structured_card(paper_id: str, user_id: str = "default"): class RelatedWorkRequest(BaseModel): - user_id: str = "default" track_id: Optional[int] = None topic: str = Field(..., min_length=1) paper_ids: Optional[List[str]] = None @@ -4025,9 +4046,9 @@ class RelatedWorkResponse(BaseModel): @router.post("/research/papers/related-work", response_model=RelatedWorkResponse) -def generate_related_work(req: RelatedWorkRequest): - items = _get_research_store().list_saved_papers( - user_id=req.user_id, track_id=req.track_id, limit=req.limit +def generate_related_work(req: RelatedWorkRequest, user_id: str = Depends(get_user_id)): + items = _research_store.list_saved_papers( + user_id=user_id, track_id=req.track_id, limit=req.limit ) papers = [item["paper"] for item in items if item.get("paper")] From 41a243bbc16b534a55bde5e09ead4b0829719a1d Mon Sep 17 00:00:00 2001 From: WenjingWang Date: Thu, 12 Mar 2026 16:43:28 +0100 Subject: [PATCH 2/9] fix(api): address PR #372 security and runtime bugs - Fix IDOR: always assign JWT-derived user_id, ignore request body user_id in gen_code, memory, and repro_context routes - Fix AttributeError: replace bare _research_store/_memory_store with _get_research_store()/_get_memory_store() in research.py; replace _store with _get_store() in memory.py - Fix missing import: add wrap_generator to gen_code.py, add sse_response to repro_context.py - Remove unused StreamingResponse import from harvest.py - Remove redundant effective_user_id variable in research.py --- src/paperbot/api/routes/gen_code.py | 5 ++--- src/paperbot/api/routes/harvest.py | 1 - src/paperbot/api/routes/memory.py | 4 ++-- src/paperbot/api/routes/repro_context.py | 8 +++----- src/paperbot/api/routes/research.py | 13 ++++++------- 5 files changed, 13 insertions(+), 18 deletions(-) diff --git a/src/paperbot/api/routes/gen_code.py b/src/paperbot/api/routes/gen_code.py index e6cb767b..b22848ce 100644 --- a/src/paperbot/api/routes/gen_code.py +++ b/src/paperbot/api/routes/gen_code.py @@ -14,7 +14,7 @@ from paperbot.core.abstractions import AgentRunContext from paperbot.api.auth.dependencies import get_user_id -from ..streaming import StreamEvent, sse_response +from ..streaming import StreamEvent, wrap_generator router = APIRouter() @@ -196,8 +196,7 @@ async def generate_code( event_log = getattr(http_request.app.state, "event_log", None) run_id = new_run_id() trace_id = new_trace_id() - # Respect authenticated user id where available; fall back to request.user_id - request.user_id = request.user_id or user_id + request.user_id = user_id return StreamingResponse( wrap_generator( diff --git a/src/paperbot/api/routes/harvest.py b/src/paperbot/api/routes/harvest.py index 5ae94a7f..7912a1d5 100644 --- a/src/paperbot/api/routes/harvest.py +++ b/src/paperbot/api/routes/harvest.py @@ -14,7 +14,6 @@ from typing import Any, Dict, List, Optional from fastapi import APIRouter, Depends, HTTPException, Query, Request -from fastapi.responses import StreamingResponse from pydantic import BaseModel, Field from paperbot.api.streaming import StreamEvent, sse_response diff --git a/src/paperbot/api/routes/memory.py b/src/paperbot/api/routes/memory.py index 65b5efd1..e147bf04 100644 --- a/src/paperbot/api/routes/memory.py +++ b/src/paperbot/api/routes/memory.py @@ -176,8 +176,8 @@ class ContextResponse(BaseModel): @router.post("/memory/context", response_model=ContextResponse) def memory_context(req: ContextRequest, user_id: str = Depends(get_user_id)): - effective_user_id = req.user_id or user_id - items = _store.search_memories( + effective_user_id = user_id + items = _get_store().search_memories( user_id=effective_user_id, workspace_id=req.workspace_id, query=req.query, diff --git a/src/paperbot/api/routes/repro_context.py b/src/paperbot/api/routes/repro_context.py index 7afbc172..0a05f90d 100644 --- a/src/paperbot/api/routes/repro_context.py +++ b/src/paperbot/api/routes/repro_context.py @@ -20,7 +20,7 @@ from fastapi.responses import StreamingResponse from pydantic import BaseModel -from paperbot.api.streaming import StreamEvent, wrap_generator +from paperbot.api.streaming import StreamEvent, wrap_generator, sse_response from paperbot.api.auth.dependencies import get_user_id from paperbot.application.services.p2c.models import ( GenerateContextRequest as P2CRequest, @@ -338,12 +338,10 @@ async def generate_context_pack( ): """Generate a P2C context pack for the given paper. Returns SSE stream.""" trace_id = set_trace_id() - # Prefer authenticated user id when available - effective_user_id = request.user_id or user_id - request.user_id = effective_user_id + request.user_id = user_id Logger.info( - f"[M2] generate_request trace_id={trace_id} paper_id={request.paper_id} user_id={effective_user_id}", + f"[M2] generate_request trace_id={trace_id} paper_id={request.paper_id} user_id={user_id}", file=LogFiles.API, ) return sse_response(_generate_stream(request), workflow="p2c_generate") diff --git a/src/paperbot/api/routes/research.py b/src/paperbot/api/routes/research.py index 5c14411d..4769d885 100644 --- a/src/paperbot/api/routes/research.py +++ b/src/paperbot/api/routes/research.py @@ -443,9 +443,8 @@ def create_track( background_tasks: BackgroundTasks, user_id: str = Depends(get_user_id), ): - effective_user_id = user_id - track = _research_store.create_track( - user_id=effective_user_id, + track = _get_research_store().create_track( + user_id=user_id, name=req.name, description=req.description, keywords=req.keywords, @@ -577,7 +576,7 @@ def get_deadline_radar( @router.get("/research/tracks/active", response_model=TrackResponse) def get_active_track(user_id: str = Depends(get_user_id)): - track = _research_store.get_active_track(user_id=user_id) + track = _get_research_store().get_active_track(user_id=user_id) if not track: raise HTTPException(status_code=404, detail="No active track for user") return TrackResponse(track=track) @@ -629,7 +628,7 @@ def activate_track( background_tasks: BackgroundTasks, user_id: str = Depends(get_user_id), ): - track = _research_store.activate_track(user_id=user_id, track_id=track_id) + track = _get_research_store().activate_track(user_id=user_id, track_id=track_id) if not track: raise HTTPException(status_code=404, detail="Track not found") _schedule_embedding_precompute(background_tasks, user_id=user_id, track_ids=[track_id]) @@ -656,7 +655,7 @@ def add_task( background_tasks: BackgroundTasks, user_id: str = Depends(get_user_id), ): - task = _research_store.add_task( + task = _get_research_store().add_task( user_id=user_id, track_id=track_id, title=req.title, @@ -740,7 +739,7 @@ def create_memory_item( scope_id=scope_id, status=req.status, ) - created, _, rows = _memory_store.add_memories(user_id=user_id, memories=[cand]) + created, _, rows = _get_memory_store().add_memories(user_id=user_id, memories=[cand]) if created <= 0 or not rows: raise HTTPException( status_code=409, detail="Duplicate memory item (same scope/kind/content)" From a96be5bf82c7402dda354d33bfb526109ecb7069 Mon Sep 17 00:00:00 2001 From: WenjingWang Date: Thu, 12 Mar 2026 18:11:22 +0100 Subject: [PATCH 3/9] fix(api): fix remaining lazy-getter and ownership-check issues per PR #372 review - repro_context: add Depends(get_user_id) + ownership checks to all context-pack routes - research: replace all bare _memory_store/_research_store/_track_router accesses with lazy getters; fix get_track_context to use Depends(get_user_id) --- src/paperbot/api/routes/repro_context.py | 24 ++++++++---- src/paperbot/api/routes/research.py | 50 ++++++++++++------------ 2 files changed, 42 insertions(+), 32 deletions(-) diff --git a/src/paperbot/api/routes/repro_context.py b/src/paperbot/api/routes/repro_context.py index 0a05f90d..25411a22 100644 --- a/src/paperbot/api/routes/repro_context.py +++ b/src/paperbot/api/routes/repro_context.py @@ -353,11 +353,11 @@ async def generate_context_pack( @router.get("") async def list_context_packs( - user_id: str = "default", paper_id: Optional[str] = None, project_id: Optional[str] = None, limit: int = 20, offset: int = 0, + user_id: str = Depends(get_user_id), ): """List context packs for a user, with optional filters.""" set_trace_id() @@ -381,7 +381,7 @@ async def list_context_packs( # ------------------------------------------------------------------ # @router.get("/{pack_id}") -async def get_context_pack(pack_id: str): +async def get_context_pack(pack_id: str, user_id: str = Depends(get_user_id)): """Return the full context pack detail.""" set_trace_id() Logger.info(f"[M2] get_pack pack_id={pack_id}", file=LogFiles.API) @@ -389,17 +389,22 @@ async def get_context_pack(pack_id: str): if pack is None: Logger.warning(f"[M2] pack_not_found pack_id={pack_id}", file=LogFiles.API) raise HTTPException(status_code=404, detail="Context pack not found.") + if pack.get("user_id") != user_id: + raise HTTPException(status_code=403, detail="Access denied.") return pack @router.get("/{pack_id}/observation/{observation_id}") -async def get_observation_detail(pack_id: str, observation_id: str): +async def get_observation_detail(pack_id: str, observation_id: str, user_id: str = Depends(get_user_id)): """Return a single observation detail by ID.""" set_trace_id() Logger.info( f"[M2] get_observation pack_id={pack_id} observation_id={observation_id}", file=LogFiles.API, ) + pack = await asyncio.to_thread(_get_store().get, pack_id) + if pack is None or pack.get("user_id") != user_id: + raise HTTPException(status_code=404, detail="Context pack not found.") observation = await asyncio.to_thread(_get_store().get_observation, pack_id, observation_id) if observation is None: Logger.warning( @@ -415,7 +420,7 @@ async def get_observation_detail(pack_id: str, observation_id: str): # ------------------------------------------------------------------ # @router.post("/{pack_id}/session") -async def create_repro_session(pack_id: str, request: CreateSessionRequest): +async def create_repro_session(pack_id: str, request: CreateSessionRequest, user_id: str = Depends(get_user_id)): """ Convert a context pack into a runbook session. @@ -425,6 +430,8 @@ async def create_repro_session(pack_id: str, request: CreateSessionRequest): if pack is None: Logger.warning(f"[M2] session_pack_not_found pack_id={pack_id}", file=LogFiles.API) raise HTTPException(status_code=404, detail="Context pack not found.") + if pack.get("user_id") != user_id: + raise HTTPException(status_code=403, detail="Access denied.") Logger.info( f"[M2] create_session pack_id={pack_id} executor={request.executor_preference}", file=LogFiles.API, @@ -460,12 +467,15 @@ async def create_repro_session(pack_id: str, request: CreateSessionRequest): # ------------------------------------------------------------------ # @router.delete("/{pack_id}") -async def delete_context_pack(pack_id: str): +async def delete_context_pack(pack_id: str, user_id: str = Depends(get_user_id)): """Soft-delete a context pack.""" set_trace_id() Logger.info(f"[M2] delete_pack pack_id={pack_id}", file=LogFiles.API) - deleted = await asyncio.to_thread(_get_store().soft_delete, pack_id) - if not deleted: + pack = await asyncio.to_thread(_get_store().get, pack_id) + if pack is None: Logger.warning(f"[M2] delete_pack_not_found pack_id={pack_id}", file=LogFiles.API) raise HTTPException(status_code=404, detail="Context pack not found.") + if pack.get("user_id") != user_id: + raise HTTPException(status_code=403, detail="Access denied.") + await asyncio.to_thread(_get_store().soft_delete, pack_id) return {"status": "deleted", "context_pack_id": pack_id} diff --git a/src/paperbot/api/routes/research.py b/src/paperbot/api/routes/research.py index 4769d885..845303b9 100644 --- a/src/paperbot/api/routes/research.py +++ b/src/paperbot/api/routes/research.py @@ -583,7 +583,7 @@ def get_active_track(user_id: str = Depends(get_user_id)): @router.get("/research/tracks/{track_id}/context", response_model=TrackContextResponse) -def get_track_context(track_id: int, user_id: str = "default"): +def get_track_context(track_id: int, user_id: str = Depends(get_user_id)): snapshot = _build_track_context_service().get_track_context( user_id=user_id, track_id=track_id, @@ -842,7 +842,7 @@ def suggest_memories( ) for m in extracted ] - created, skipped, rows = _memory_store.add_memories(user_id=user_id, memories=pending) + created, skipped, rows = _get_memory_store().add_memories(user_id=user_id, memories=pending) if scope_type == "track": _schedule_embedding_precompute( background_tasks, user_id=user_id, track_ids=[int(scope_id or 0)] @@ -870,7 +870,7 @@ def moderate_memory_item( req: MemoryModerateRequest, user_id: str = Depends(get_user_id), ): - updated = _memory_store.update_item( + updated = _get_memory_store().update_item( user_id=user_id, item_id=item_id, status=req.status, @@ -1104,7 +1104,7 @@ def precompute_embeddings( req: PrecomputeEmbeddingsRequest, user_id: str = Depends(get_user_id), ): - result = _track_router.precompute_track_embeddings( + result = _get_track_router().precompute_track_embeddings( user_id=user_id, track_ids=req.track_ids or None ) return PrecomputeEmbeddingsResponse(user_id=user_id, result=result) @@ -1229,7 +1229,7 @@ def add_paper_feedback( Logger.warning(f"Failed to save paper to library: {e}", file=LogFiles.HARVEST) Logger.info("Recording paper feedback to research store", file=LogFiles.HARVEST) - fb = _research_store.add_paper_feedback( + fb = _get_research_store().add_paper_feedback( user_id=user_id, track_id=track_id, paper_id=req.paper_id, # Always use external ID for consistency @@ -1398,7 +1398,7 @@ def update_paper_status( req: PaperReadingStatusRequest, user_id: str = Depends(get_user_id), ): - status = _research_store.set_paper_reading_status( + status = _get_research_store().set_paper_reading_status( user_id=user_id, paper_id=paper_id, status=req.status, @@ -1682,7 +1682,7 @@ async def discover_from_seed(req: DiscoverySeedRequest, user_id: str = Depends(g @router.post("/research/collections", response_model=PaperCollectionResponse) def create_collection(req: PaperCollectionCreateRequest, user_id: str = Depends(get_user_id)): try: - collection = _research_store.create_collection( + collection = _get_research_store().create_collection( user_id=user_id, name=req.name, description=req.description, @@ -1712,7 +1712,7 @@ def list_collections( @router.patch("/research/collections/{collection_id}", response_model=PaperCollectionResponse) def update_collection(collection_id: int, req: PaperCollectionUpdateRequest, user_id: str = Depends(get_user_id)): try: - collection = _research_store.update_collection( + collection = _get_research_store().update_collection( user_id=user_id, collection_id=collection_id, name=req.name, @@ -1748,7 +1748,7 @@ def list_collection_items( response_model=PaperCollectionItemsResponse, ) def upsert_collection_item(collection_id: int, req: PaperCollectionItemUpsertRequest, user_id: str = Depends(get_user_id)): - item = _research_store.upsert_collection_item( + item = _get_research_store().upsert_collection_item( user_id=user_id, collection_id=collection_id, paper_id=req.paper_id, @@ -1757,7 +1757,7 @@ def upsert_collection_item(collection_id: int, req: PaperCollectionItemUpsertReq ) if item is None: raise HTTPException(status_code=404, detail="Collection or paper not found") - items = _research_store.list_collection_items(user_id=user_id, collection_id=collection_id) + items = _get_research_store().list_collection_items(user_id=user_id, collection_id=collection_id) return PaperCollectionItemsResponse( user_id=user_id, collection_id=collection_id, items=items ) @@ -1768,7 +1768,7 @@ def upsert_collection_item(collection_id: int, req: PaperCollectionItemUpsertReq response_model=PaperCollectionItemsResponse, ) def patch_collection_item(collection_id: int, paper_id: str, req: PaperCollectionItemPatchRequest, user_id: str = Depends(get_user_id)): - item = _research_store.upsert_collection_item( + item = _get_research_store().upsert_collection_item( user_id=user_id, collection_id=collection_id, paper_id=paper_id, @@ -1777,7 +1777,7 @@ def patch_collection_item(collection_id: int, paper_id: str, req: PaperCollectio ) if item is None: raise HTTPException(status_code=404, detail="Collection or paper not found") - items = _research_store.list_collection_items(user_id=user_id, collection_id=collection_id) + items = _get_research_store().list_collection_items(user_id=user_id, collection_id=collection_id) return PaperCollectionItemsResponse( user_id=user_id, collection_id=collection_id, items=items ) @@ -1785,7 +1785,7 @@ def patch_collection_item(collection_id: int, paper_id: str, req: PaperCollectio @router.delete("/research/collections/{collection_id}/items/{paper_id}") def delete_collection_item(collection_id: int, paper_id: str, user_id: str = Depends(get_user_id)): - ok = _research_store.remove_collection_item( + ok = _get_research_store().remove_collection_item( user_id=user_id, collection_id=collection_id, paper_id=paper_id, @@ -1908,7 +1908,7 @@ def import_bibtex(req: BibtexImportRequest, user_id: str = Depends(get_user_id)) track_pk = int(track["id"]) paper_store = _get_paper_store() - existing_saved_ids = _research_store.list_paper_feedback_ids( + existing_saved_ids = _get_research_store().list_paper_feedback_ids( user_id=user_id, track_id=track_pk, action="save", @@ -1949,7 +1949,7 @@ def import_bibtex(req: BibtexImportRequest, user_id: str = Depends(get_user_id)) "citation_key": str(entry.get("key") or ""), "entry_type": str(entry.get("entry_type") or ""), } - _research_store.add_paper_feedback( + _get_research_store().add_paper_feedback( user_id=user_id, track_id=track_pk, paper_id=paper_ref, @@ -2020,7 +2020,7 @@ def pull_from_zotero(req: ZoteroSyncRequest, user_id: str = Depends(get_user_id) raise HTTPException(status_code=502, detail=f"Failed to pull from Zotero: {exc}") from exc paper_store = _get_paper_store() - existing_saved_ids = _research_store.list_paper_feedback_ids( + existing_saved_ids = _get_research_store().list_paper_feedback_ids( user_id=user_id, track_id=track_pk, action="save", @@ -2061,7 +2061,7 @@ def pull_from_zotero(req: ZoteroSyncRequest, user_id: str = Depends(get_user_id) "zotero_library_type": req.library_type, "zotero_library_id": req.library_id, } - _research_store.add_paper_feedback( + _get_research_store().add_paper_feedback( user_id=user_id, track_id=track_pk, paper_id=paper_ref, @@ -2108,12 +2108,12 @@ def push_to_zotero(req: ZoteroPushRequest, user_id: str = Depends(get_user_id)): from paperbot.infrastructure.connectors.zotero_connector import ZoteroConnector if req.track_id is not None: - track = _research_store.get_track(user_id=user_id, track_id=req.track_id) + track = _get_research_store().get_track(user_id=user_id, track_id=req.track_id) if not track: raise HTTPException(status_code=404, detail="Track not found") connector = ZoteroConnector() - local_items = _research_store.list_saved_papers( + local_items = _get_research_store().list_saved_papers( user_id=user_id, track_id=req.track_id, sort_by="saved_at", @@ -2226,7 +2226,7 @@ def discover_track_anchors( def set_anchor_action(track_id: int, author_id: int, req: AnchorActionRequest, user_id: str = Depends(get_user_id)): _ensure_anchor_feature_enabled() - track = _research_store.get_track(user_id=user_id, track_id=track_id) + track = _get_research_store().get_track(user_id=user_id, track_id=track_id) if not track: raise HTTPException(status_code=404, detail="Track not found") @@ -2272,7 +2272,7 @@ def list_anchor_actions(track_id: int, user_id: str = Depends(get_user_id)): @router.get("/research/papers/{paper_id}", response_model=PaperDetailResponse) def get_paper_detail(paper_id: str, user_id: str = Depends(get_user_id)): - detail = _research_store.get_paper_detail(paper_id=paper_id, user_id=user_id) + detail = _get_research_store().get_paper_detail(paper_id=paper_id, user_id=user_id) if not detail: raise HTTPException(status_code=404, detail="Paper not found in registry") return PaperDetailResponse(detail=detail) @@ -2296,7 +2296,7 @@ class RouterSuggestResponse(BaseModel): @router.post("/research/router/suggest", response_model=RouterSuggestResponse) def suggest_track(req: RouterSuggestRequest, user_id: str = Depends(get_user_id)): - active = _research_store.get_active_track(user_id=user_id) + active = _get_research_store().get_active_track(user_id=user_id) if not active: return RouterSuggestResponse(suggestion=None) grounded_query = _get_workflow_query_grounder().ground_query( @@ -2411,7 +2411,7 @@ async def build_context(req: ContextRequest, user_id: str = Depends(get_user_id) if req.activate_track_id is not None: Logger.info("Activating research track", file=LogFiles.HARVEST) - activated = _research_store.activate_track( + activated = _get_research_store().activate_track( user_id=user_id, track_id=req.activate_track_id ) if not activated: @@ -3982,7 +3982,7 @@ class StructuredCardResponse(BaseModel): @router.get("/research/papers/{paper_id}/card", response_model=StructuredCardResponse) def get_structured_card(paper_id: str, user_id: str = Depends(get_user_id)): - detail = _research_store.get_paper_detail(paper_id=paper_id, user_id=user_id) + detail = _get_research_store().get_paper_detail(paper_id=paper_id, user_id=user_id) if not detail: raise HTTPException(status_code=404, detail="Paper not found") @@ -4046,7 +4046,7 @@ class RelatedWorkResponse(BaseModel): @router.post("/research/papers/related-work", response_model=RelatedWorkResponse) def generate_related_work(req: RelatedWorkRequest, user_id: str = Depends(get_user_id)): - items = _research_store.list_saved_papers( + items = _get_research_store().list_saved_papers( user_id=user_id, track_id=req.track_id, limit=req.limit ) papers = [item["paper"] for item in items if item.get("paper")] From a80412a7361353d3fb1481b9d704461a54397d67 Mon Sep 17 00:00:00 2001 From: WenjingWang Date: Thu, 12 Mar 2026 19:29:57 +0100 Subject: [PATCH 4/9] fix(api): enforce strict auth on all user-scoped routes and clean up code smells --- src/paperbot/api/auth/dependencies.py | 14 ++++ src/paperbot/api/routes/gen_code.py | 11 ++- src/paperbot/api/routes/memory.py | 17 +++-- src/paperbot/api/routes/repro_context.py | 14 ++-- src/paperbot/api/routes/research.py | 92 ++++++++++++------------ 5 files changed, 80 insertions(+), 68 deletions(-) diff --git a/src/paperbot/api/auth/dependencies.py b/src/paperbot/api/auth/dependencies.py index 4bcaaa94..657513c5 100644 --- a/src/paperbot/api/auth/dependencies.py +++ b/src/paperbot/api/auth/dependencies.py @@ -73,3 +73,17 @@ def get_user_id( if user is None: return "default" return str(user.id) + + +def get_required_user_id( + credentials: Optional[HTTPAuthorizationCredentials] = Depends(bearer), +) -> str: + """Like get_user_id but rejects unauthenticated requests even when AUTH_OPTIONAL=true. + + Use this for any endpoint that writes or reads user-scoped data, so that + anonymous callers cannot touch the shared "default" namespace. + """ + user = _resolve_user(credentials) + if user is None: + raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Authentication required") + return str(user.id) diff --git a/src/paperbot/api/routes/gen_code.py b/src/paperbot/api/routes/gen_code.py index b22848ce..b38b2281 100644 --- a/src/paperbot/api/routes/gen_code.py +++ b/src/paperbot/api/routes/gen_code.py @@ -12,7 +12,7 @@ from paperbot.application.collaboration.message_schema import new_run_id, new_trace_id from paperbot.core.abstractions import AgentRunContext -from paperbot.api.auth.dependencies import get_user_id +from paperbot.api.auth.dependencies import get_required_user_id from ..streaming import StreamEvent, wrap_generator @@ -30,7 +30,7 @@ class GenCodeRequest(BaseModel): async def gen_code_stream( - request: GenCodeRequest, *, event_log=None, run_id: str = "", trace_id: str = "" + request: GenCodeRequest, *, user_id: str, event_log=None, run_id: str = "", trace_id: str = "" ): """Stream code generation progress""" try: @@ -96,7 +96,7 @@ async def gen_code_stream( result = await agent.reproduce_from_paper( paper_context, output_dir=output_dir, - user_id=request.user_id, + user_id=user_id, event_log=event_log, run_id=run_id, trace_id=trace_id, @@ -186,7 +186,7 @@ async def gen_code_stream( async def generate_code( request: GenCodeRequest, http_request: Request, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): """ Generate code from paper and stream progress. @@ -196,11 +196,10 @@ async def generate_code( event_log = getattr(http_request.app.state, "event_log", None) run_id = new_run_id() trace_id = new_trace_id() - request.user_id = user_id return StreamingResponse( wrap_generator( - gen_code_stream(request, event_log=event_log, run_id=run_id, trace_id=trace_id), + gen_code_stream(request, user_id=user_id, event_log=event_log, run_id=run_id, trace_id=trace_id), workflow="gen_code", run_id=run_id, trace_id=trace_id, diff --git a/src/paperbot/api/routes/memory.py b/src/paperbot/api/routes/memory.py index e147bf04..365db97f 100644 --- a/src/paperbot/api/routes/memory.py +++ b/src/paperbot/api/routes/memory.py @@ -7,7 +7,7 @@ from pydantic import BaseModel, Field from paperbot.infrastructure.stores.memory_store import SqlAlchemyMemoryStore -from paperbot.api.auth.dependencies import get_user_id +from paperbot.api.auth.dependencies import get_required_user_id from paperbot.memory import build_memory_context, extract_memories, parse_chat_log from paperbot.memory.schema import MemoryCandidate @@ -57,7 +57,7 @@ class IngestResponse(BaseModel): @router.post("/memory/ingest", response_model=IngestResponse) async def ingest_memory( file: UploadFile = File(...), - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), workspace_id: Optional[str] = Query(None, description="Optional workspace/project namespace."), platform: Optional[str] = Query(None, description="Hint: chatgpt/gemini/claude/..."), use_llm: bool = Query(False, description="Use configured LLM to extract memories (falls back on heuristics)."), @@ -123,7 +123,7 @@ class MemoryListResponse(BaseModel): @router.get("/memory/list", response_model=MemoryListResponse) def list_memories( - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), limit: int = 100, kind: Optional[str] = None, workspace_id: Optional[str] = None, @@ -175,10 +175,9 @@ class ContextResponse(BaseModel): @router.post("/memory/context", response_model=ContextResponse) -def memory_context(req: ContextRequest, user_id: str = Depends(get_user_id)): - effective_user_id = user_id +def memory_context(req: ContextRequest, user_id: str = Depends(get_required_user_id)): items = _get_store().search_memories( - user_id=effective_user_id, + user_id=user_id, workspace_id=req.workspace_id, query=req.query, limit=req.limit, @@ -197,7 +196,7 @@ def memory_context(req: ContextRequest, user_id: str = Depends(get_user_id)): ] ctx = build_memory_context(cands, max_items=req.limit) return ContextResponse( - user_id=effective_user_id, + user_id=user_id, query=req.query, context=ctx, items=[ @@ -229,7 +228,7 @@ class MemoryItemUpdateRequest(BaseModel): @router.patch("/memory/items/{item_id}", response_model=MemoryItemOut) -def update_memory_item(user_id: str, item_id: int, body: MemoryItemUpdateRequest): +def update_memory_item(item_id: int, body: MemoryItemUpdateRequest, user_id: str = Depends(get_required_user_id)): updated = _get_store().update_item( user_id=user_id, item_id=item_id, @@ -259,11 +258,11 @@ def update_memory_item(user_id: str, item_id: int, body: MemoryItemUpdateRequest @router.delete("/memory/items/{item_id}") def delete_memory_item( - user_id: str, item_id: int, actor_id: str = "system", reason: str = "", hard: bool = False, + user_id: str = Depends(get_required_user_id), ): if hard: ok = _get_store().hard_delete_item(user_id=user_id, item_id=item_id, actor_id=actor_id) diff --git a/src/paperbot/api/routes/repro_context.py b/src/paperbot/api/routes/repro_context.py index 25411a22..21250def 100644 --- a/src/paperbot/api/routes/repro_context.py +++ b/src/paperbot/api/routes/repro_context.py @@ -21,7 +21,7 @@ from pydantic import BaseModel from paperbot.api.streaming import StreamEvent, wrap_generator, sse_response -from paperbot.api.auth.dependencies import get_user_id +from paperbot.api.auth.dependencies import get_user_id, get_required_user_id from paperbot.application.services.p2c.models import ( GenerateContextRequest as P2CRequest, RawPaperData, @@ -334,7 +334,7 @@ async def _write_paper_scope_memories( @router.post("/generate") async def generate_context_pack( request: GenerateContextPackRequest, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): """Generate a P2C context pack for the given paper. Returns SSE stream.""" trace_id = set_trace_id() @@ -357,7 +357,7 @@ async def list_context_packs( project_id: Optional[str] = None, limit: int = 20, offset: int = 0, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): """List context packs for a user, with optional filters.""" set_trace_id() @@ -381,7 +381,7 @@ async def list_context_packs( # ------------------------------------------------------------------ # @router.get("/{pack_id}") -async def get_context_pack(pack_id: str, user_id: str = Depends(get_user_id)): +async def get_context_pack(pack_id: str, user_id: str = Depends(get_required_user_id)): """Return the full context pack detail.""" set_trace_id() Logger.info(f"[M2] get_pack pack_id={pack_id}", file=LogFiles.API) @@ -395,7 +395,7 @@ async def get_context_pack(pack_id: str, user_id: str = Depends(get_user_id)): @router.get("/{pack_id}/observation/{observation_id}") -async def get_observation_detail(pack_id: str, observation_id: str, user_id: str = Depends(get_user_id)): +async def get_observation_detail(pack_id: str, observation_id: str, user_id: str = Depends(get_required_user_id)): """Return a single observation detail by ID.""" set_trace_id() Logger.info( @@ -420,7 +420,7 @@ async def get_observation_detail(pack_id: str, observation_id: str, user_id: str # ------------------------------------------------------------------ # @router.post("/{pack_id}/session") -async def create_repro_session(pack_id: str, request: CreateSessionRequest, user_id: str = Depends(get_user_id)): +async def create_repro_session(pack_id: str, request: CreateSessionRequest, user_id: str = Depends(get_required_user_id)): """ Convert a context pack into a runbook session. @@ -467,7 +467,7 @@ async def create_repro_session(pack_id: str, request: CreateSessionRequest, user # ------------------------------------------------------------------ # @router.delete("/{pack_id}") -async def delete_context_pack(pack_id: str, user_id: str = Depends(get_user_id)): +async def delete_context_pack(pack_id: str, user_id: str = Depends(get_required_user_id)): """Soft-delete a context pack.""" set_trace_id() Logger.info(f"[M2] delete_pack pack_id={pack_id}", file=LogFiles.API) diff --git a/src/paperbot/api/routes/research.py b/src/paperbot/api/routes/research.py index 845303b9..de1a11df 100644 --- a/src/paperbot/api/routes/research.py +++ b/src/paperbot/api/routes/research.py @@ -36,7 +36,7 @@ from paperbot.infrastructure.stores.memory_store import SqlAlchemyMemoryStore from paperbot.infrastructure.stores.research_store import SqlAlchemyResearchStore from paperbot.infrastructure.stores.workflow_metric_store import WorkflowMetricStore -from paperbot.api.auth.dependencies import get_user_id +from paperbot.api.auth.dependencies import get_user_id, get_required_user_id from paperbot.memory.eval.collector import MemoryMetricCollector from paperbot.memory.extractor import extract_memories from paperbot.memory.schema import MemoryCandidate, NormalizedMessage @@ -441,7 +441,7 @@ def _serialize_track_context_response( def create_track( req: TrackCreateRequest, background_tasks: BackgroundTasks, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): track = _get_research_store().create_track( user_id=user_id, @@ -468,7 +468,7 @@ class TrackListResponse(BaseModel): @router.get("/research/tracks", response_model=TrackListResponse) def list_tracks( - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), include_archived: bool = Query(False), limit: int = Query(100, ge=1, le=500), ): @@ -486,7 +486,7 @@ class DeadlineRadarResponse(BaseModel): @router.get("/research/deadlines/radar", response_model=DeadlineRadarResponse) def get_deadline_radar( - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), days: int = Query(180, ge=7, le=365), ccf_levels: str = Query("A,B,C"), field: Optional[str] = None, @@ -575,7 +575,7 @@ def get_deadline_radar( @router.get("/research/tracks/active", response_model=TrackResponse) -def get_active_track(user_id: str = Depends(get_user_id)): +def get_active_track(user_id: str = Depends(get_required_user_id)): track = _get_research_store().get_active_track(user_id=user_id) if not track: raise HTTPException(status_code=404, detail="No active track for user") @@ -583,7 +583,7 @@ def get_active_track(user_id: str = Depends(get_user_id)): @router.get("/research/tracks/{track_id}/context", response_model=TrackContextResponse) -def get_track_context(track_id: int, user_id: str = Depends(get_user_id)): +def get_track_context(track_id: int, user_id: str = Depends(get_required_user_id)): snapshot = _build_track_context_service().get_track_context( user_id=user_id, track_id=track_id, @@ -598,7 +598,7 @@ def update_track( track_id: int, req: TrackUpdateRequest, background_tasks: BackgroundTasks, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): update_data = req.model_dump(exclude_unset=True, exclude_none=True) @@ -626,7 +626,7 @@ def update_track( def activate_track( track_id: int, background_tasks: BackgroundTasks, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): track = _get_research_store().activate_track(user_id=user_id, track_id=track_id) if not track: @@ -653,7 +653,7 @@ def add_task( track_id: int, req: TaskCreateRequest, background_tasks: BackgroundTasks, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): task = _get_research_store().add_task( user_id=user_id, @@ -680,7 +680,7 @@ class TaskListResponse(BaseModel): @router.get("/research/tracks/{track_id}/tasks", response_model=TaskListResponse) def list_tasks( track_id: int, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), status: Optional[str] = None, limit: int = Query(100, ge=1, le=500), ): @@ -719,7 +719,7 @@ def _resolve_track_scope_id( def create_memory_item( req: MemoryItemCreateRequest, background_tasks: BackgroundTasks, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): scope_type = (req.scope_type or "global").strip() or "global" scope_id = _resolve_track_scope_id(user_id, scope_type, req.scope_id) @@ -758,7 +758,7 @@ class MemoryItemListResponse(BaseModel): @router.get("/research/memory/items", response_model=MemoryItemListResponse) def list_memory_items( - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), scope_type: Optional[str] = None, scope_id: Optional[str] = None, kind: Optional[str] = None, @@ -780,7 +780,7 @@ def list_memory_items( @router.get("/research/memory/inbox", response_model=MemoryItemListResponse) def list_memory_inbox( - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), track_id: Optional[int] = None, limit: int = Query(100, ge=1, le=500), ): @@ -815,7 +815,7 @@ class MemorySuggestResponse(BaseModel): def suggest_memories( req: MemorySuggestRequest, background_tasks: BackgroundTasks, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): scope_type = (req.scope_type or "global").strip() or "global" scope_id = _resolve_track_scope_id(user_id, scope_type, req.scope_id) @@ -868,7 +868,7 @@ class MemoryModerateRequest(BaseModel): def moderate_memory_item( item_id: int, req: MemoryModerateRequest, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): updated = _get_memory_store().update_item( user_id=user_id, @@ -900,7 +900,7 @@ class BulkModerateResponse(BaseModel): def bulk_moderate( req: BulkModerateRequest, background_tasks: BackgroundTasks, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): result = _build_track_memory_service().bulk_moderate( user_id=user_id, @@ -953,7 +953,7 @@ class BulkMoveResponse(BaseModel): def bulk_move( req: BulkMoveRequest, background_tasks: BackgroundTasks, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): try: result = _build_track_memory_service().bulk_move( @@ -998,7 +998,7 @@ class MemoryFeedbackResponse(BaseModel): @router.post("/research/memory/feedback", response_model=MemoryFeedbackResponse) def record_memory_feedback( req: MemoryFeedbackRequest, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): """ Record user feedback on retrieved memories. @@ -1056,7 +1056,7 @@ class ClearTrackMemoryResponse(BaseModel): def clear_track_memory( track_id: int, background_tasks: BackgroundTasks, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), confirm: bool = Query(False), ): if not confirm: @@ -1102,7 +1102,7 @@ class PrecomputeEmbeddingsResponse(BaseModel): @router.post("/research/embeddings/precompute", response_model=PrecomputeEmbeddingsResponse) def precompute_embeddings( req: PrecomputeEmbeddingsRequest, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): result = _get_track_router().precompute_track_embeddings( user_id=user_id, track_ids=req.track_ids or None @@ -1118,7 +1118,7 @@ class EvalSummaryResponse(BaseModel): @router.get("/research/evals/summary", response_model=EvalSummaryResponse) def eval_summary( - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), track_id: Optional[int] = None, days: int = Query(30, ge=1, le=365), ): @@ -1155,7 +1155,7 @@ class PaperFeedbackResponse(BaseModel): def add_paper_feedback( req: PaperFeedbackRequest, background_tasks: BackgroundTasks, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): set_trace_id() # Initialize trace_id for this request Logger.info(f"Received paper feedback request, action={req.action}", file=LogFiles.HARVEST) @@ -1266,7 +1266,7 @@ class PaperFeedbackListResponse(BaseModel): @router.get("/research/tracks/{track_id}/papers/feedback", response_model=PaperFeedbackListResponse) def list_paper_feedback( track_id: int, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), action: Optional[str] = None, limit: int = Query(200, ge=1, le=1000), ): @@ -1396,7 +1396,7 @@ class PaperRepoListResponse(BaseModel): def update_paper_status( paper_id: str, req: PaperReadingStatusRequest, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): status = _get_research_store().set_paper_reading_status( user_id=user_id, @@ -1412,7 +1412,7 @@ def update_paper_status( @router.get("/research/papers/saved", response_model=SavedPapersResponse) def list_saved_papers( - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), track_id: Optional[int] = None, collection_id: Optional[int] = None, sort_by: str = Query("saved_at"), @@ -1429,7 +1429,7 @@ def list_saved_papers( @router.post("/research/discovery/seed", response_model=DiscoverySeedResponse) -async def discover_from_seed(req: DiscoverySeedRequest, user_id: str = Depends(get_user_id)): +async def discover_from_seed(req: DiscoverySeedRequest, user_id: str = Depends(get_required_user_id)): if req.year_from and req.year_to and req.year_from > req.year_to: raise HTTPException(status_code=400, detail="year_from must be <= year_to") @@ -1680,7 +1680,7 @@ async def discover_from_seed(req: DiscoverySeedRequest, user_id: str = Depends(g @router.post("/research/collections", response_model=PaperCollectionResponse) -def create_collection(req: PaperCollectionCreateRequest, user_id: str = Depends(get_user_id)): +def create_collection(req: PaperCollectionCreateRequest, user_id: str = Depends(get_required_user_id)): try: collection = _get_research_store().create_collection( user_id=user_id, @@ -1695,7 +1695,7 @@ def create_collection(req: PaperCollectionCreateRequest, user_id: str = Depends( @router.get("/research/collections", response_model=PaperCollectionListResponse) def list_collections( - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), include_archived: bool = Query(False), track_id: Optional[int] = Query(default=None), limit: int = Query(200, ge=1, le=1000), @@ -1710,7 +1710,7 @@ def list_collections( @router.patch("/research/collections/{collection_id}", response_model=PaperCollectionResponse) -def update_collection(collection_id: int, req: PaperCollectionUpdateRequest, user_id: str = Depends(get_user_id)): +def update_collection(collection_id: int, req: PaperCollectionUpdateRequest, user_id: str = Depends(get_required_user_id)): try: collection = _get_research_store().update_collection( user_id=user_id, @@ -1732,7 +1732,7 @@ def update_collection(collection_id: int, req: PaperCollectionUpdateRequest, use ) def list_collection_items( collection_id: int, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), limit: int = Query(500, ge=1, le=5000), ): items = _get_research_store().list_collection_items( @@ -1747,7 +1747,7 @@ def list_collection_items( "/research/collections/{collection_id}/items", response_model=PaperCollectionItemsResponse, ) -def upsert_collection_item(collection_id: int, req: PaperCollectionItemUpsertRequest, user_id: str = Depends(get_user_id)): +def upsert_collection_item(collection_id: int, req: PaperCollectionItemUpsertRequest, user_id: str = Depends(get_required_user_id)): item = _get_research_store().upsert_collection_item( user_id=user_id, collection_id=collection_id, @@ -1767,7 +1767,7 @@ def upsert_collection_item(collection_id: int, req: PaperCollectionItemUpsertReq "/research/collections/{collection_id}/items/{paper_id}", response_model=PaperCollectionItemsResponse, ) -def patch_collection_item(collection_id: int, paper_id: str, req: PaperCollectionItemPatchRequest, user_id: str = Depends(get_user_id)): +def patch_collection_item(collection_id: int, paper_id: str, req: PaperCollectionItemPatchRequest, user_id: str = Depends(get_required_user_id)): item = _get_research_store().upsert_collection_item( user_id=user_id, collection_id=collection_id, @@ -1784,7 +1784,7 @@ def patch_collection_item(collection_id: int, paper_id: str, req: PaperCollectio @router.delete("/research/collections/{collection_id}/items/{paper_id}") -def delete_collection_item(collection_id: int, paper_id: str, user_id: str = Depends(get_user_id)): +def delete_collection_item(collection_id: int, paper_id: str, user_id: str = Depends(get_required_user_id)): ok = _get_research_store().remove_collection_item( user_id=user_id, collection_id=collection_id, @@ -1798,7 +1798,7 @@ def delete_collection_item(collection_id: int, paper_id: str, user_id: str = Dep @router.get("/research/tracks/{track_id}/feed", response_model=TrackFeedResponse) def get_track_feed( track_id: int, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), limit: int = Query(20, ge=1, le=100), offset: int = Query(0, ge=0), ): @@ -1824,7 +1824,7 @@ def get_track_feed( @router.get("/research/papers/export") def export_papers( - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), track_id: Optional[int] = None, format: str = Query("bibtex", pattern="^(bibtex|ris|markdown|csl_json)$"), ): @@ -1894,7 +1894,7 @@ class BibtexImportResponse(BaseModel): @router.post("/research/papers/import/bibtex", response_model=BibtexImportResponse) -def import_bibtex(req: BibtexImportRequest, user_id: str = Depends(get_user_id)): +def import_bibtex(req: BibtexImportRequest, user_id: str = Depends(get_required_user_id)): entries = _parse_bibtex_entries(req.content) if not entries: raise HTTPException(status_code=400, detail="No valid BibTeX entries found") @@ -1997,7 +1997,7 @@ class ZoteroPullResponse(BaseModel): @router.post("/research/integrations/zotero/pull", response_model=ZoteroPullResponse) -def pull_from_zotero(req: ZoteroSyncRequest, user_id: str = Depends(get_user_id)): +def pull_from_zotero(req: ZoteroSyncRequest, user_id: str = Depends(get_required_user_id)): from paperbot.infrastructure.connectors.zotero_connector import ZoteroConnector track = _resolve_or_create_import_track( @@ -2104,7 +2104,7 @@ class ZoteroPushResponse(BaseModel): @router.post("/research/integrations/zotero/push", response_model=ZoteroPushResponse) -def push_to_zotero(req: ZoteroPushRequest, user_id: str = Depends(get_user_id)): +def push_to_zotero(req: ZoteroPushRequest, user_id: str = Depends(get_required_user_id)): from paperbot.infrastructure.connectors.zotero_connector import ZoteroConnector if req.track_id is not None: @@ -2186,7 +2186,7 @@ def push_to_zotero(req: ZoteroPushRequest, user_id: str = Depends(get_user_id)): @router.get("/research/tracks/{track_id}/anchors/discover", response_model=AnchorDiscoverResponse) def discover_track_anchors( track_id: int, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), limit: int = Query(20, ge=1, le=100), window_days: int = Query(365, ge=30, le=3650), personalized: bool = Query(True), @@ -2223,7 +2223,7 @@ def discover_track_anchors( "/research/tracks/{track_id}/anchors/{author_id}/action", response_model=AnchorActionResponse, ) -def set_anchor_action(track_id: int, author_id: int, req: AnchorActionRequest, user_id: str = Depends(get_user_id)): +def set_anchor_action(track_id: int, author_id: int, req: AnchorActionRequest, user_id: str = Depends(get_required_user_id)): _ensure_anchor_feature_enabled() track = _get_research_store().get_track(user_id=user_id, track_id=track_id) @@ -2259,7 +2259,7 @@ def set_anchor_action(track_id: int, author_id: int, req: AnchorActionRequest, u @router.get("/research/tracks/{track_id}/anchors/actions", response_model=AnchorActionListResponse) -def list_anchor_actions(track_id: int, user_id: str = Depends(get_user_id)): +def list_anchor_actions(track_id: int, user_id: str = Depends(get_required_user_id)): _ensure_anchor_feature_enabled() track = _get_research_store().get_track(user_id=user_id, track_id=track_id) @@ -2271,7 +2271,7 @@ def list_anchor_actions(track_id: int, user_id: str = Depends(get_user_id)): @router.get("/research/papers/{paper_id}", response_model=PaperDetailResponse) -def get_paper_detail(paper_id: str, user_id: str = Depends(get_user_id)): +def get_paper_detail(paper_id: str, user_id: str = Depends(get_required_user_id)): detail = _get_research_store().get_paper_detail(paper_id=paper_id, user_id=user_id) if not detail: raise HTTPException(status_code=404, detail="Paper not found in registry") @@ -2295,7 +2295,7 @@ class RouterSuggestResponse(BaseModel): @router.post("/research/router/suggest", response_model=RouterSuggestResponse) -def suggest_track(req: RouterSuggestRequest, user_id: str = Depends(get_user_id)): +def suggest_track(req: RouterSuggestRequest, user_id: str = Depends(get_required_user_id)): active = _get_research_store().get_active_track(user_id=user_id) if not active: return RouterSuggestResponse(suggestion=None) @@ -2399,7 +2399,7 @@ def get_evidence_coverage( @router.post("/research/context", response_model=ContextResponse) -async def build_context(req: ContextRequest, user_id: str = Depends(get_user_id)): +async def build_context(req: ContextRequest, user_id: str = Depends(get_required_user_id)): set_trace_id() # Initialize trace_id for this request Logger.info("Received build context request", file=LogFiles.HARVEST) @@ -3981,7 +3981,7 @@ class StructuredCardResponse(BaseModel): @router.get("/research/papers/{paper_id}/card", response_model=StructuredCardResponse) -def get_structured_card(paper_id: str, user_id: str = Depends(get_user_id)): +def get_structured_card(paper_id: str, user_id: str = Depends(get_required_user_id)): detail = _get_research_store().get_paper_detail(paper_id=paper_id, user_id=user_id) if not detail: raise HTTPException(status_code=404, detail="Paper not found") @@ -4045,7 +4045,7 @@ class RelatedWorkResponse(BaseModel): @router.post("/research/papers/related-work", response_model=RelatedWorkResponse) -def generate_related_work(req: RelatedWorkRequest, user_id: str = Depends(get_user_id)): +def generate_related_work(req: RelatedWorkRequest, user_id: str = Depends(get_required_user_id)): items = _get_research_store().list_saved_papers( user_id=user_id, track_id=req.track_id, limit=req.limit ) From 44e9fb1213d387994bc217b1fd85491d83190224 Mon Sep 17 00:00:00 2001 From: WenjingWang Date: Fri, 13 Mar 2026 00:20:01 +0100 Subject: [PATCH 5/9] fix(auth): extend JWT-based identity to remaining routes. Related to #151 --- src/paperbot/api/routes/chat.py | 8 ++- src/paperbot/api/routes/gen_code.py | 21 ++---- src/paperbot/api/routes/harvest.py | 8 +-- src/paperbot/api/routes/intelligence.py | 35 ++++------ src/paperbot/api/routes/memory.py | 1 - src/paperbot/api/routes/repro_context.py | 14 ++-- src/paperbot/api/routes/research.py | 2 +- .../application/services/identity_resolver.py | 10 ++- .../infrastructure/stores/paper_store.py | 1 + tests/unit/test_intelligence_routes.py | 67 +++++++++++++------ 10 files changed, 92 insertions(+), 75 deletions(-) diff --git a/src/paperbot/api/routes/chat.py b/src/paperbot/api/routes/chat.py index 8ba1073b..fca4e393 100644 --- a/src/paperbot/api/routes/chat.py +++ b/src/paperbot/api/routes/chat.py @@ -4,7 +4,9 @@ from typing import List, Optional -from fastapi import APIRouter +from fastapi import APIRouter, Depends + +from paperbot.api.auth.dependencies import get_required_user_id from pydantic import BaseModel, Field, field_validator from ...application.collaboration.message_schema import new_trace_id @@ -128,11 +130,13 @@ async def chat_stream(request: ChatRequest, *, trace_id: str): @router.post("/chat") -async def chat(request: ChatRequest): +async def chat(request: ChatRequest, user_id: str = Depends(get_required_user_id)): """ Chat with PaperBot AI and stream response. Returns Server-Sent Events with streaming text. """ trace_id = new_trace_id() + # Override any client-provided user_id with authenticated user + request.user_id = user_id return sse_response(chat_stream(request, trace_id=trace_id), workflow="chat", trace_id=trace_id) diff --git a/src/paperbot/api/routes/gen_code.py b/src/paperbot/api/routes/gen_code.py index b38b2281..ac785790 100644 --- a/src/paperbot/api/routes/gen_code.py +++ b/src/paperbot/api/routes/gen_code.py @@ -7,20 +7,18 @@ from typing import Optional from fastapi import APIRouter, Depends, Request -from fastapi.responses import StreamingResponse from pydantic import BaseModel from paperbot.application.collaboration.message_schema import new_run_id, new_trace_id from paperbot.core.abstractions import AgentRunContext from paperbot.api.auth.dependencies import get_required_user_id -from ..streaming import StreamEvent, wrap_generator +from ..streaming import StreamEvent, sse_response router = APIRouter() class GenCodeRequest(BaseModel): - user_id: str = "default" title: str abstract: str method_section: Optional[str] = None @@ -197,16 +195,9 @@ async def generate_code( run_id = new_run_id() trace_id = new_trace_id() - return StreamingResponse( - wrap_generator( - gen_code_stream(request, user_id=user_id, event_log=event_log, run_id=run_id, trace_id=trace_id), - workflow="gen_code", - run_id=run_id, - trace_id=trace_id, - ), - media_type="text/event-stream", - headers={ - "Cache-Control": "no-cache", - "Connection": "keep-alive", - }, + return sse_response( + gen_code_stream(request, user_id=user_id, event_log=event_log, run_id=run_id, trace_id=trace_id), + workflow="gen_code", + run_id=run_id, + trace_id=trace_id, ) diff --git a/src/paperbot/api/routes/harvest.py b/src/paperbot/api/routes/harvest.py index 7912a1d5..1ecaaa7f 100644 --- a/src/paperbot/api/routes/harvest.py +++ b/src/paperbot/api/routes/harvest.py @@ -23,7 +23,7 @@ HarvestPipeline, HarvestProgress, ) -from paperbot.api.auth.dependencies import get_user_id +from paperbot.api.auth.dependencies import get_required_user_id from paperbot.infrastructure.stores.paper_store import PaperStore, paper_to_dict from paperbot.utils.logging_config import LogFiles, Logger, clear_trace_id, set_trace_id @@ -327,7 +327,7 @@ class LibraryResponse(BaseModel): @router.get("/papers/library", response_model=LibraryResponse) def get_user_library( - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), track_id: Optional[int] = Query(None, description="Filter by track"), actions: Optional[str] = Query(None, description="Filter by actions (comma-separated)"), sort_by: str = Query("saved_at", description="Sort field"), @@ -397,7 +397,7 @@ class SavePaperRequest(BaseModel): def save_paper_to_library( paper_id: int, request: SavePaperRequest, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): """ Save a paper to user's library. @@ -425,7 +425,7 @@ def save_paper_to_library( @router.delete("/papers/{paper_id}/save") def remove_paper_from_library( paper_id: int, - user_id: str = Depends(get_user_id), + user_id: str = Depends(get_required_user_id), ): """Remove a paper from user's library.""" store = _get_paper_store() diff --git a/src/paperbot/api/routes/intelligence.py b/src/paperbot/api/routes/intelligence.py index b5a6c729..74de7471 100644 --- a/src/paperbot/api/routes/intelligence.py +++ b/src/paperbot/api/routes/intelligence.py @@ -3,17 +3,16 @@ import re from typing import Any, Dict, List, Optional -from fastapi import APIRouter, BackgroundTasks, HTTPException, Query +from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query from pydantic import BaseModel, Field from paperbot.infrastructure.services.intelligence_radar_service import IntelligenceRadarService from paperbot.infrastructure.stores.research_store import SqlAlchemyResearchStore +from paperbot.api.auth.dependencies import get_required_user_id router = APIRouter() _service: Optional[IntelligenceRadarService] = None _research_store = SqlAlchemyResearchStore() -_DEFAULT_USER_ID = "default" - _SIGNAL_STOPWORDS = { "about", "across", @@ -97,20 +96,11 @@ class IntelligenceFeedResponse(BaseModel): subreddits: List[str] = Field(default_factory=list) -def _resolve_feed_user_id(requested_user_id: Optional[str]) -> str: - user_id = str(requested_user_id or _DEFAULT_USER_ID).strip() or _DEFAULT_USER_ID - if user_id != _DEFAULT_USER_ID: - raise HTTPException( - status_code=403, - detail="cross-user intelligence access requires authenticated user context", - ) - return _DEFAULT_USER_ID - @router.get("/intelligence/feed", response_model=IntelligenceFeedResponse) def get_intelligence_feed( background_tasks: BackgroundTasks, - user_id: str = Query("default"), + user_id: str = Depends(get_required_user_id), limit: int = Query(6, ge=1, le=20), refresh: bool = Query(False), source: Optional[str] = Query(None), @@ -123,22 +113,21 @@ def get_intelligence_feed( sort_order: str = Query("desc", pattern="^(asc|desc)$"), track_id: Optional[int] = Query(None, ge=1), ): - resolved_user_id = _resolve_feed_user_id(user_id) service = _get_service() refresh_scheduled = False if refresh: - service.refresh(user_id=resolved_user_id) - elif service.needs_refresh(user_id=resolved_user_id): - cached = service.list_feed(user_id=resolved_user_id, limit=1) + service.refresh(user_id=user_id) + elif service.needs_refresh(user_id=user_id): + cached = service.list_feed(user_id=user_id, limit=1) if cached: - background_tasks.add_task(service.refresh, user_id=resolved_user_id) + background_tasks.add_task(service.refresh, user_id=user_id) refresh_scheduled = True else: - service.refresh(user_id=resolved_user_id) + service.refresh(user_id=user_id) rows = service.list_feed( - user_id=resolved_user_id, + user_id=user_id, limit=max(int(limit) * 10, 50), source=source, keyword=keyword, @@ -146,7 +135,7 @@ def get_intelligence_feed( sort_by=sort_by, sort_order=sort_order, ) - annotated_rows = [_annotate_intelligence_row(user_id=resolved_user_id, row=row) for row in rows] + annotated_rows = [_annotate_intelligence_row(user_id=user_id, row=row) for row in rows] if track_id: annotated_rows = [ row @@ -154,11 +143,11 @@ def get_intelligence_feed( if any(int(track.get("track_id") or 0) == int(track_id) for track in row.get("matched_tracks") or []) ] - profile = service.build_profile(user_id=resolved_user_id) + profile = service.build_profile(user_id=user_id) return IntelligenceFeedResponse( items=[_to_response_item(row) for row in annotated_rows[: max(1, int(limit))]], - refreshed_at=service.latest_refresh(user_id=resolved_user_id), + refreshed_at=service.latest_refresh(user_id=user_id), refresh_scheduled=refresh_scheduled, keywords=profile.keywords, watch_repos=profile.watch_repos, diff --git a/src/paperbot/api/routes/memory.py b/src/paperbot/api/routes/memory.py index 365db97f..6ef08043 100644 --- a/src/paperbot/api/routes/memory.py +++ b/src/paperbot/api/routes/memory.py @@ -160,7 +160,6 @@ def list_memories( class ContextRequest(BaseModel): - user_id: str = "default" workspace_id: Optional[str] = None query: str = Field(..., min_length=1) limit: int = 8 diff --git a/src/paperbot/api/routes/repro_context.py b/src/paperbot/api/routes/repro_context.py index 21250def..42b9f460 100644 --- a/src/paperbot/api/routes/repro_context.py +++ b/src/paperbot/api/routes/repro_context.py @@ -21,7 +21,7 @@ from pydantic import BaseModel from paperbot.api.streaming import StreamEvent, wrap_generator, sse_response -from paperbot.api.auth.dependencies import get_user_id, get_required_user_id +from paperbot.api.auth.dependencies import get_required_user_id from paperbot.application.services.p2c.models import ( GenerateContextRequest as P2CRequest, RawPaperData, @@ -51,7 +51,6 @@ def _get_store() -> SqlAlchemyReproContextStore: class GenerateContextPackRequest(BaseModel): paper_id: str - user_id: str = "default" project_id: Optional[str] = None track_id: Optional[int] = None depth: Literal["fast", "standard", "deep"] = "standard" @@ -70,7 +69,7 @@ class CreateSessionRequest(BaseModel): # POST /generate (SSE) # # ------------------------------------------------------------------ # -async def _generate_stream(request: GenerateContextPackRequest): +async def _generate_stream(request: GenerateContextPackRequest, user_id: str): """SSE generator for context pack generation via Module 1 ExtractionOrchestrator.""" pack_id = new_context_pack_id() Logger.info( @@ -83,7 +82,7 @@ async def _generate_stream(request: GenerateContextPackRequest): await asyncio.to_thread( _get_store().save, pack_id=pack_id, - user_id=request.user_id, + user_id=user_id, paper_id=request.paper_id, depth=request.depth, pack_data={}, @@ -169,7 +168,7 @@ async def on_stage_complete(stage_name: str, observations: list, warnings: list) p2c_request = P2CRequest( paper_id=request.paper_id, - user_id=request.user_id, + user_id=user_id, project_id=request.project_id, track_id=request.track_id, depth=request.depth, @@ -265,7 +264,7 @@ async def _run() -> None: # Write observations to paper-scope memory so future P2C runs can reuse them. await _write_paper_scope_memories( paper_id=request.paper_id, - user_id=request.user_id, + user_id=user_id, observations=pack.observations, ) Logger.info( @@ -338,13 +337,12 @@ async def generate_context_pack( ): """Generate a P2C context pack for the given paper. Returns SSE stream.""" trace_id = set_trace_id() - request.user_id = user_id Logger.info( f"[M2] generate_request trace_id={trace_id} paper_id={request.paper_id} user_id={user_id}", file=LogFiles.API, ) - return sse_response(_generate_stream(request), workflow="p2c_generate") + return sse_response(_generate_stream(request, user_id), workflow="p2c_generate") # ------------------------------------------------------------------ # diff --git a/src/paperbot/api/routes/research.py b/src/paperbot/api/routes/research.py index de1a11df..463ad0d9 100644 --- a/src/paperbot/api/routes/research.py +++ b/src/paperbot/api/routes/research.py @@ -36,7 +36,7 @@ from paperbot.infrastructure.stores.memory_store import SqlAlchemyMemoryStore from paperbot.infrastructure.stores.research_store import SqlAlchemyResearchStore from paperbot.infrastructure.stores.workflow_metric_store import WorkflowMetricStore -from paperbot.api.auth.dependencies import get_user_id, get_required_user_id +from paperbot.api.auth.dependencies import get_required_user_id from paperbot.memory.eval.collector import MemoryMetricCollector from paperbot.memory.extractor import extract_memories from paperbot.memory.schema import MemoryCandidate, NormalizedMessage diff --git a/src/paperbot/application/services/identity_resolver.py b/src/paperbot/application/services/identity_resolver.py index fd3e807c..1bfd7345 100644 --- a/src/paperbot/application/services/identity_resolver.py +++ b/src/paperbot/application/services/identity_resolver.py @@ -49,7 +49,15 @@ def resolve( pid = (external_id or "").strip() hints = hints or {} - # 0. Numeric ID → direct lookup + # 0a. library_paper_id hint → already-resolved internal ID, use directly + lib_id = hints.get("library_paper_id") + if lib_id is not None: + try: + return int(lib_id) + except (TypeError, ValueError): + pass + + # 0b. Numeric ID → direct lookup if pid.isdigit(): with self._provider.session() as session: row = session.execute( diff --git a/src/paperbot/infrastructure/stores/paper_store.py b/src/paperbot/infrastructure/stores/paper_store.py index 5edaad05..8fba734a 100644 --- a/src/paperbot/infrastructure/stores/paper_store.py +++ b/src/paperbot/infrastructure/stores/paper_store.py @@ -575,6 +575,7 @@ def _create_model(self, paper: HarvestedPaper, now: datetime) -> PaperModel: fields_of_study_json=json.dumps(paper.fields_of_study, ensure_ascii=False), primary_source=paper.source.value, sources_json=json.dumps([paper.source.value], ensure_ascii=False), + first_seen_at=now, created_at=now, updated_at=now, ) diff --git a/tests/unit/test_intelligence_routes.py b/tests/unit/test_intelligence_routes.py index df1b82ff..efc52781 100644 --- a/tests/unit/test_intelligence_routes.py +++ b/tests/unit/test_intelligence_routes.py @@ -1,6 +1,7 @@ from fastapi.testclient import TestClient from paperbot.api import main as api_main +from paperbot.api.auth import dependencies as auth_deps from paperbot.api.routes import intelligence as intelligence_route from paperbot.infrastructure.services.intelligence_radar_service import RadarProfile @@ -89,25 +90,36 @@ def list_tracks(self, *, user_id: str, include_archived: bool, limit: int): ] +def _override_user_id(user_id: str): + def _dep_override(): + return user_id + + return _dep_override + + def test_intelligence_feed_route_returns_external_signal_payload(monkeypatch): service = _FakeIntelligenceService() monkeypatch.setattr(intelligence_route, "_service", service) monkeypatch.setattr(intelligence_route, "_research_store", _FakeResearchStore()) - with TestClient(api_main.app) as client: - resp = client.get( - "/api/intelligence/feed", - params={ - "user_id": "default", - "limit": 1, - "source": "reddit", - "keyword": "rag", - "repo": "org/rag-agent", - "sort_by": "delta", - "sort_order": "desc", - "track_id": 7, - }, - ) + app = api_main.app + app.dependency_overrides[auth_deps.get_required_user_id] = _override_user_id("default") + try: + with TestClient(app) as client: + resp = client.get( + "/api/intelligence/feed", + params={ + "limit": 1, + "source": "reddit", + "keyword": "rag", + "repo": "org/rag-agent", + "sort_by": "delta", + "sort_order": "desc", + "track_id": 7, + }, + ) + finally: + app.dependency_overrides.clear() assert resp.status_code == 200 payload = resp.json() @@ -152,14 +164,29 @@ def test_intelligence_feed_route_returns_external_signal_payload(monkeypatch): assert "rag" in item["research_query"] -def test_intelligence_feed_route_rejects_cross_user_access(monkeypatch): +def test_intelligence_feed_route_uses_authenticated_user_id(monkeypatch): service = _FakeIntelligenceService() monkeypatch.setattr(intelligence_route, "_service", service) monkeypatch.setattr(intelligence_route, "_research_store", _FakeResearchStore()) - with TestClient(api_main.app) as client: - resp = client.get("/api/intelligence/feed", params={"user_id": "u-radar"}) + app = api_main.app + app.dependency_overrides[auth_deps.get_required_user_id] = _override_user_id("u-radar") + try: + with TestClient(app) as client: + resp = client.get( + "/api/intelligence/feed", + params={ + "limit": 1, + "source": "reddit", + "keyword": "rag", + "repo": "org/rag-agent", + "sort_by": "delta", + "sort_order": "desc", + "track_id": 7, + }, + ) + finally: + app.dependency_overrides.clear() - assert resp.status_code == 403 - assert resp.json()["detail"] == "cross-user intelligence access requires authenticated user context" - assert service.list_feed_calls == [] + assert resp.status_code == 200 + assert service.list_feed_calls[0]["user_id"] == "u-radar" From b709034141cf1229f9250915ecb0f82ea17bfdde Mon Sep 17 00:00:00 2001 From: WenjingWang Date: Fri, 13 Mar 2026 00:20:20 +0100 Subject: [PATCH 6/9] feat(web): add NextAuth authentication and frontend auth flows. Related to #151 --- web/next.config.ts | 32 ++ web/package-lock.json | 94 +++++ web/package.json | 3 +- web/src/app/api/_utils/auth-headers.ts | 38 ++ web/src/app/api/auth/[...nextauth]/route.ts | 3 + web/src/app/api/auth/forgot-password/route.ts | 13 + web/src/app/api/auth/login-check/route.ts | 19 + .../app/api/auth/me/change-password/route.ts | 14 + web/src/app/api/auth/me/route.ts | 27 ++ web/src/app/api/auth/register/route.ts | 13 + web/src/app/api/auth/reset-password/route.ts | 13 + web/src/app/api/gen-code/route.ts | 7 +- web/src/app/api/research/_base.ts | 12 +- .../api/research/paperscool/daily/route.ts | 5 +- .../app/api/research/repro/context/route.ts | 5 +- web/src/app/api/runbook/delete/route.ts | 5 +- .../app/api/runbook/revert-project/route.ts | 5 +- web/src/app/api/runbook/smoke/route.ts | 5 +- web/src/app/api/studio/chat/route.ts | 6 +- web/src/app/dashboard/page.tsx | 17 +- web/src/app/forgot-password/page.tsx | 115 ++++++ web/src/app/layout.tsx | 5 +- web/src/app/login/page.tsx | 210 ++++++++++ web/src/app/papers/[id]/page.tsx | 5 +- web/src/app/papers/page.tsx | 10 +- web/src/app/register/page.tsx | 316 ++++++++++++++ web/src/app/research/page.tsx | 9 +- web/src/app/reset-password/page.tsx | 190 +++++++++ web/src/app/settings/page.tsx | 387 ++++++++++++++++-- web/src/auth.ts | 93 +++++ web/src/components/dashboard/ActivityFeed.tsx | 7 +- .../dashboard/DashboardReadingQueuePanel.tsx | 7 +- .../dashboard/TrackSpotlightSection.tsx | 12 +- web/src/components/layout/Sidebar.tsx | 90 +++- .../research/DiscoveryGraphWorkspace.tsx | 5 +- web/src/components/research/FeedTab.tsx | 8 +- web/src/components/research/MemoryTab.tsx | 8 +- .../components/research/ResearchDashboard.tsx | 24 +- .../research/ResearchDiscoveryPage.tsx | 9 +- .../components/research/ResearchPageNew.tsx | 23 +- .../components/research/SavedPapersList.tsx | 22 +- web/src/components/research/SavedTab.tsx | 16 +- .../components/scholars/ScholarsWatchlist.tsx | 3 +- web/src/hooks/useContextPackGeneration.ts | 1 - web/src/lib/api.ts | 49 ++- web/src/lib/config.ts | 2 + web/src/lib/dashboard-api.ts | 47 +-- web/src/middleware.ts | 36 ++ web/src/types/next-auth.d.ts | 9 + 49 files changed, 1848 insertions(+), 206 deletions(-) create mode 100644 web/src/app/api/_utils/auth-headers.ts create mode 100644 web/src/app/api/auth/[...nextauth]/route.ts create mode 100644 web/src/app/api/auth/forgot-password/route.ts create mode 100644 web/src/app/api/auth/login-check/route.ts create mode 100644 web/src/app/api/auth/me/change-password/route.ts create mode 100644 web/src/app/api/auth/me/route.ts create mode 100644 web/src/app/api/auth/register/route.ts create mode 100644 web/src/app/api/auth/reset-password/route.ts create mode 100644 web/src/app/forgot-password/page.tsx create mode 100644 web/src/app/login/page.tsx create mode 100644 web/src/app/register/page.tsx create mode 100644 web/src/app/reset-password/page.tsx create mode 100644 web/src/auth.ts create mode 100644 web/src/lib/config.ts create mode 100644 web/src/middleware.ts create mode 100644 web/src/types/next-auth.d.ts diff --git a/web/next.config.ts b/web/next.config.ts index 8e168edc..433b6006 100644 --- a/web/next.config.ts +++ b/web/next.config.ts @@ -6,7 +6,39 @@ const nextConfig: NextConfig = { keepAlive: true, }, async rewrites() { + // Important: Next route handlers under /app/api take precedence over rewrites. + // This list only applies to paths that do NOT have a corresponding file-based + // route. We still keep explicit "bypass" rules here for clarity. return [ + // Keep NextAuth handlers on the Next.js side + { + source: '/api/auth/:path*', + destination: '/api/auth/:path*', + }, + // Keep our own proxy/utility routes handled by Next (app/api/**) + // NOTE: app/api route files already win over rewrites; these entries are + // mainly defensive and for future explicit exceptions. + { + source: '/api/research/:path*', + destination: '/api/research/:path*', + }, + { + source: '/api/runbook/:path*', + destination: '/api/runbook/:path*', + }, + { + source: '/api/studio/:path*', + destination: '/api/studio/:path*', + }, + { + source: '/api/sandbox/:path*', + destination: '/api/sandbox/:path*', + }, + { + source: '/api/papers/:path*', + destination: '/api/papers/:path*', + }, + // Default: proxy any other /api/* calls to the FastAPI backend { source: '/api/:path*', destination: 'http://localhost:8000/api/:path*', diff --git a/web/package-lock.json b/web/package-lock.json index ffd4e34b..8743f3ed 100644 --- a/web/package-lock.json +++ b/web/package-lock.json @@ -37,6 +37,7 @@ "framer-motion": "^12.23.26", "lucide-react": "^0.562.0", "next": "16.1.0", + "next-auth": "5.0.0-beta.30", "next-themes": "^0.4.6", "radix-ui": "^1.4.3", "react": "19.2.3", @@ -197,6 +198,35 @@ "url": "https://github.com/sponsors/sindresorhus" } }, + "node_modules/@auth/core": { + "version": "0.41.0", + "resolved": "https://registry.npmjs.org/@auth/core/-/core-0.41.0.tgz", + "integrity": "sha512-Wd7mHPQ/8zy6Qj7f4T46vg3aoor8fskJm6g2Zyj064oQ3+p0xNZXAV60ww0hY+MbTesfu29kK14Zk5d5JTazXQ==", + "license": "ISC", + "dependencies": { + "@panva/hkdf": "^1.2.1", + "jose": "^6.0.6", + "oauth4webapi": "^3.3.0", + "preact": "10.24.3", + "preact-render-to-string": "6.5.11" + }, + "peerDependencies": { + "@simplewebauthn/browser": "^9.0.1", + "@simplewebauthn/server": "^9.0.2", + "nodemailer": "^6.8.0" + }, + "peerDependenciesMeta": { + "@simplewebauthn/browser": { + "optional": true + }, + "@simplewebauthn/server": { + "optional": true + }, + "nodemailer": { + "optional": true + } + } + }, "node_modules/@babel/code-frame": { "version": "7.27.1", "resolved": "https://registry.npmjs.org/@babel/code-frame/-/code-frame-7.27.1.tgz", @@ -1921,6 +1951,15 @@ "node": ">=8.0.0" } }, + "node_modules/@panva/hkdf": { + "version": "1.2.1", + "resolved": "https://registry.npmjs.org/@panva/hkdf/-/hkdf-1.2.1.tgz", + "integrity": "sha512-6oclG6Y3PiDFcoyk8srjLfVKyMfVCKJ27JwNPViuXziFpmdz+MZnZN/aKY0JGXgYuO/VghU0jcOAZgWXZ1Dmrw==", + "license": "MIT", + "funding": { + "url": "https://github.com/sponsors/panva" + } + }, "node_modules/@playwright/test": { "version": "1.58.2", "resolved": "https://registry.npmjs.org/@playwright/test/-/test-1.58.2.tgz", @@ -11934,6 +11973,33 @@ } } }, + "node_modules/next-auth": { + "version": "5.0.0-beta.30", + "resolved": "https://registry.npmjs.org/next-auth/-/next-auth-5.0.0-beta.30.tgz", + "integrity": "sha512-+c51gquM3F6nMVmoAusRJ7RIoY0K4Ts9HCCwyy/BRoe4mp3msZpOzYMyb5LAYc1wSo74PMQkGDcaghIO7W6Xjg==", + "license": "ISC", + "dependencies": { + "@auth/core": "0.41.0" + }, + "peerDependencies": { + "@simplewebauthn/browser": "^9.0.1", + "@simplewebauthn/server": "^9.0.2", + "next": "^14.0.0-0 || ^15.0.0 || ^16.0.0", + "nodemailer": "^7.0.7", + "react": "^18.2.0 || ^19.0.0" + }, + "peerDependenciesMeta": { + "@simplewebauthn/browser": { + "optional": true + }, + "@simplewebauthn/server": { + "optional": true + }, + "nodemailer": { + "optional": true + } + } + }, "node_modules/next-themes": { "version": "0.4.6", "resolved": "https://registry.npmjs.org/next-themes/-/next-themes-0.4.6.tgz", @@ -11979,6 +12045,15 @@ "dev": true, "license": "MIT" }, + "node_modules/oauth4webapi": { + "version": "3.8.5", + "resolved": "https://registry.npmjs.org/oauth4webapi/-/oauth4webapi-3.8.5.tgz", + "integrity": "sha512-A8jmyUckVhRJj5lspguklcl90Ydqk61H3dcU0oLhH3Yv13KpAliKTt5hknpGGPZSSfOwGyraNEFmofDYH+1kSg==", + "license": "MIT", + "funding": { + "url": "https://github.com/sponsors/panva" + } + }, "node_modules/object-assign": { "version": "4.1.1", "resolved": "https://registry.npmjs.org/object-assign/-/object-assign-4.1.1.tgz", @@ -12388,6 +12463,25 @@ "node": "^10 || ^12 || >=14" } }, + "node_modules/preact": { + "version": "10.24.3", + "resolved": "https://registry.npmjs.org/preact/-/preact-10.24.3.tgz", + "integrity": "sha512-Z2dPnBnMUfyQfSQ+GBdsGa16hz35YmLmtTLhM169uW944hYL6xzTYkJjC07j+Wosz733pMWx0fgON3JNw1jJQA==", + "license": "MIT", + "funding": { + "type": "opencollective", + "url": "https://opencollective.com/preact" + } + }, + "node_modules/preact-render-to-string": { + "version": "6.5.11", + "resolved": "https://registry.npmjs.org/preact-render-to-string/-/preact-render-to-string-6.5.11.tgz", + "integrity": "sha512-ubnauqoGczeGISiOh6RjX0/cdaF8v/oDXIjO85XALCQjwQP+SB4RDXXtvZ6yTYSjG+PC1QRP2AhPgCEsM2EvUw==", + "license": "MIT", + "peerDependencies": { + "preact": ">=10" + } + }, "node_modules/prelude-ls": { "version": "1.2.1", "resolved": "https://registry.npmjs.org/prelude-ls/-/prelude-ls-1.2.1.tgz", diff --git a/web/package.json b/web/package.json index 6ced83f6..eb37ad43 100644 --- a/web/package.json +++ b/web/package.json @@ -55,7 +55,8 @@ "xterm": "^5.3.0", "xterm-addon-fit": "^0.8.0", "zustand": "^5.0.9", - "@radix-ui/react-popover": "^1.1.15" + "@radix-ui/react-popover": "^1.1.15", + "next-auth": "5.0.0-beta.30" }, "devDependencies": { "@playwright/test": "^1.58.2", diff --git a/web/src/app/api/_utils/auth-headers.ts b/web/src/app/api/_utils/auth-headers.ts new file mode 100644 index 00000000..7d08f618 --- /dev/null +++ b/web/src/app/api/_utils/auth-headers.ts @@ -0,0 +1,38 @@ +import { auth } from "@/auth" + +export function backendBaseUrl(): string { + return process.env.BACKEND_BASE_URL || "http://127.0.0.1:8000" +} + +export async function withBackendAuth( + req: Request, + base: HeadersInit = {} +): Promise { + const headers = new Headers(base as any) + + // Prefer client-provided Authorization header if present + const incoming = req.headers.get("authorization") + if (incoming) { + headers.set("authorization", incoming) + return headers + } + + // Otherwise pull token from NextAuth session (server-side) + try { + const session = await auth() + const token = (session as any)?.accessToken as string | undefined + if (token) { + headers.set("authorization", `Bearer ${token}`) + } else { + console.warn("[auth-headers] session.accessToken is missing", { + hasSession: !!session, + userId: (session as any)?.userId, + provider: (session as any)?.provider, + }) + } + } catch (e) { + console.error("[auth-headers] auth() threw:", e) + } + return headers +} + diff --git a/web/src/app/api/auth/[...nextauth]/route.ts b/web/src/app/api/auth/[...nextauth]/route.ts new file mode 100644 index 00000000..4129ec49 --- /dev/null +++ b/web/src/app/api/auth/[...nextauth]/route.ts @@ -0,0 +1,3 @@ +import { handlers } from "@/auth" +export const { GET, POST } = handlers + diff --git a/web/src/app/api/auth/forgot-password/route.ts b/web/src/app/api/auth/forgot-password/route.ts new file mode 100644 index 00000000..7135b7de --- /dev/null +++ b/web/src/app/api/auth/forgot-password/route.ts @@ -0,0 +1,13 @@ +import { NextRequest, NextResponse } from "next/server" +import { backendBaseUrl } from "@/app/api/_utils/auth-headers" + +export async function POST(req: NextRequest) { + const body = await req.json() + const res = await fetch(`${backendBaseUrl()}/api/auth/forgot-password`, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(body), + }) + const data = await res.json().catch(() => null) + return NextResponse.json(data, { status: res.status }) +} diff --git a/web/src/app/api/auth/login-check/route.ts b/web/src/app/api/auth/login-check/route.ts new file mode 100644 index 00000000..410c4bb0 --- /dev/null +++ b/web/src/app/api/auth/login-check/route.ts @@ -0,0 +1,19 @@ +export const runtime = "nodejs" + +import { backendBaseUrl } from "../../_utils/auth-headers" + +export async function POST(req: Request) { + const body = await req.text() + const res = await fetch(`${backendBaseUrl()}/api/auth/login`, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body, + }).catch(() => null) + + if (!res) return Response.json({ detail: "Service unavailable" }, { status: 502 }) + const text = await res.text() + return new Response(text, { + status: res.status, + headers: { "Content-Type": "application/json" }, + }) +} diff --git a/web/src/app/api/auth/me/change-password/route.ts b/web/src/app/api/auth/me/change-password/route.ts new file mode 100644 index 00000000..e593324f --- /dev/null +++ b/web/src/app/api/auth/me/change-password/route.ts @@ -0,0 +1,14 @@ +import { NextRequest, NextResponse } from "next/server" +import { backendBaseUrl, withBackendAuth } from "@/app/api/_utils/auth-headers" + +export async function POST(req: NextRequest) { + const body = await req.json() + const headers = await withBackendAuth(req, { "Content-Type": "application/json" }) + const res = await fetch(`${backendBaseUrl()}/api/auth/me/change-password`, { + method: "POST", + headers, + body: JSON.stringify(body), + }) + const data = await res.json().catch(() => null) + return NextResponse.json(data, { status: res.status }) +} diff --git a/web/src/app/api/auth/me/route.ts b/web/src/app/api/auth/me/route.ts new file mode 100644 index 00000000..f74601f2 --- /dev/null +++ b/web/src/app/api/auth/me/route.ts @@ -0,0 +1,27 @@ +import { NextRequest, NextResponse } from "next/server" +import { backendBaseUrl, withBackendAuth } from "@/app/api/_utils/auth-headers" + +export async function GET(req: NextRequest) { + const headers = await withBackendAuth(req) + const res = await fetch(`${backendBaseUrl()}/api/auth/me`, { headers }) + const data = await res.json().catch(() => null) + return NextResponse.json(data, { status: res.status }) +} + +export async function PATCH(req: NextRequest) { + const body = await req.json() + const headers = await withBackendAuth(req, { "Content-Type": "application/json" }) + const res = await fetch(`${backendBaseUrl()}/api/auth/me`, { + method: "PATCH", + headers, + body: JSON.stringify(body), + }) + const data = await res.json().catch(() => null) + return NextResponse.json(data, { status: res.status }) +} + +export async function DELETE(req: NextRequest) { + const headers = await withBackendAuth(req) + const res = await fetch(`${backendBaseUrl()}/api/auth/me`, { method: "DELETE", headers }) + return new NextResponse(null, { status: res.status }) +} diff --git a/web/src/app/api/auth/register/route.ts b/web/src/app/api/auth/register/route.ts new file mode 100644 index 00000000..bfea4723 --- /dev/null +++ b/web/src/app/api/auth/register/route.ts @@ -0,0 +1,13 @@ +import { NextRequest, NextResponse } from "next/server" +import { backendBaseUrl } from "@/app/api/_utils/auth-headers" + +export async function POST(req: NextRequest) { + const body = await req.json() + const res = await fetch(`${backendBaseUrl()}/api/auth/register`, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(body), + }) + const data = await res.json().catch(() => null) + return NextResponse.json(data, { status: res.status }) +} diff --git a/web/src/app/api/auth/reset-password/route.ts b/web/src/app/api/auth/reset-password/route.ts new file mode 100644 index 00000000..25e68124 --- /dev/null +++ b/web/src/app/api/auth/reset-password/route.ts @@ -0,0 +1,13 @@ +import { NextRequest, NextResponse } from "next/server" +import { backendBaseUrl } from "@/app/api/_utils/auth-headers" + +export async function POST(req: NextRequest) { + const body = await req.json() + const res = await fetch(`${backendBaseUrl()}/api/auth/reset-password`, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(body), + }) + const data = await res.json().catch(() => null) + return NextResponse.json(data, { status: res.status }) +} diff --git a/web/src/app/api/gen-code/route.ts b/web/src/app/api/gen-code/route.ts index 7e4e78b7..c5f78144 100644 --- a/web/src/app/api/gen-code/route.ts +++ b/web/src/app/api/gen-code/route.ts @@ -4,13 +4,15 @@ function apiBaseUrl() { return process.env.PAPERBOT_API_BASE_URL || "http://127.0.0.1:8000" } +import { withBackendAuth } from "../_utils/auth-headers" + export async function POST(req: Request) { const body = await req.text() const upstream = await fetch(`${apiBaseUrl()}/api/gen-code`, { method: "POST", - headers: { + headers: await withBackendAuth(req, { "Content-Type": req.headers.get("content-type") || "application/json", - }, + }), body, }) @@ -24,4 +26,3 @@ export async function POST(req: Request) { headers, }) } - diff --git a/web/src/app/api/research/_base.ts b/web/src/app/api/research/_base.ts index 5e396e1e..fad9f1a7 100644 --- a/web/src/app/api/research/_base.ts +++ b/web/src/app/api/research/_base.ts @@ -8,12 +8,16 @@ export async function proxyJson(req: Request, upstreamUrl: string, method: strin const timeout = setTimeout(() => controller.abort(), 120_000) // 2 min timeout try { + const baseHeaders = { + method, + Accept: "application/json", + "Content-Type": req.headers.get("content-type") || "application/json", + } as Record + const { withBackendAuth } = await import("../_utils/auth-headers") + const headers = await withBackendAuth(req, baseHeaders) const upstream = await fetch(upstreamUrl, { method, - headers: { - Accept: "application/json", - "Content-Type": req.headers.get("content-type") || "application/json", - }, + headers, body, signal: controller.signal, }) diff --git a/web/src/app/api/research/paperscool/daily/route.ts b/web/src/app/api/research/paperscool/daily/route.ts index 4e288498..3cce1e0b 100644 --- a/web/src/app/api/research/paperscool/daily/route.ts +++ b/web/src/app/api/research/paperscool/daily/route.ts @@ -3,6 +3,7 @@ export const runtime = "nodejs" import { Agent } from "undici" import { apiBaseUrl } from "../../_base" +import { withBackendAuth } from "../../../_utils/auth-headers" // Keep SSE proxy streams alive during long backend phases (LLM/Judge). const sseDispatcher = new Agent({ @@ -18,10 +19,10 @@ export async function POST(req: Request) { try { upstream = await fetch(`${apiBaseUrl()}/api/research/paperscool/daily`, { method: "POST", - headers: { + headers: await withBackendAuth(req, { "Content-Type": contentType, Accept: "text/event-stream, application/json", - }, + }), body, dispatcher: sseDispatcher, } as RequestInit & { dispatcher: Agent }) diff --git a/web/src/app/api/research/repro/context/route.ts b/web/src/app/api/research/repro/context/route.ts index c268f631..5f783cd3 100644 --- a/web/src/app/api/research/repro/context/route.ts +++ b/web/src/app/api/research/repro/context/route.ts @@ -1,4 +1,5 @@ import { apiBaseUrl, proxyJson } from "../../_base" +import { withBackendAuth } from "../../../_utils/auth-headers" export const runtime = "nodejs" @@ -11,10 +12,10 @@ export async function POST(req: Request) { const body = await req.text() const upstream = await fetch(`${apiBaseUrl()}/api/research/repro/context/generate`, { method: "POST", - headers: { + headers: await withBackendAuth(req, { "Content-Type": req.headers.get("content-type") || "application/json", Accept: "text/event-stream", - }, + }), body, }) diff --git a/web/src/app/api/runbook/delete/route.ts b/web/src/app/api/runbook/delete/route.ts index 9c515b47..5f9ce689 100644 --- a/web/src/app/api/runbook/delete/route.ts +++ b/web/src/app/api/runbook/delete/route.ts @@ -8,10 +8,10 @@ export async function POST(req: Request) { const body = await req.text() const upstream = await fetch(`${apiBaseUrl()}/api/runbook/delete`, { method: "POST", - headers: { + headers: await (await import("../../_utils/auth-headers")).withBackendAuth(req, { "Content-Type": req.headers.get("content-type") || "application/json", Accept: "application/json", - }, + }), body, }) const text = await upstream.text() @@ -23,4 +23,3 @@ export async function POST(req: Request) { }, }) } - diff --git a/web/src/app/api/runbook/revert-project/route.ts b/web/src/app/api/runbook/revert-project/route.ts index c31dc6cc..a1ecd7b5 100644 --- a/web/src/app/api/runbook/revert-project/route.ts +++ b/web/src/app/api/runbook/revert-project/route.ts @@ -8,10 +8,10 @@ export async function POST(req: Request) { const body = await req.text() const upstream = await fetch(`${apiBaseUrl()}/api/runbook/revert-project`, { method: "POST", - headers: { + headers: await (await import("../../_utils/auth-headers")).withBackendAuth(req, { "Content-Type": req.headers.get("content-type") || "application/json", Accept: "application/json", - }, + }), body, }) const text = await upstream.text() @@ -23,4 +23,3 @@ export async function POST(req: Request) { }, }) } - diff --git a/web/src/app/api/runbook/smoke/route.ts b/web/src/app/api/runbook/smoke/route.ts index 4e9cf7f6..19dde6bb 100644 --- a/web/src/app/api/runbook/smoke/route.ts +++ b/web/src/app/api/runbook/smoke/route.ts @@ -8,9 +8,9 @@ export async function POST(req: Request) { const body = await req.text() const upstream = await fetch(`${apiBaseUrl()}/api/runbook/smoke`, { method: "POST", - headers: { + headers: await (await import("../../_utils/auth-headers")).withBackendAuth(req, { "Content-Type": req.headers.get("content-type") || "application/json", - }, + }), body, }) @@ -23,4 +23,3 @@ export async function POST(req: Request) { }, }) } - diff --git a/web/src/app/api/studio/chat/route.ts b/web/src/app/api/studio/chat/route.ts index 3fe12b8f..47adfd53 100644 --- a/web/src/app/api/studio/chat/route.ts +++ b/web/src/app/api/studio/chat/route.ts @@ -4,13 +4,15 @@ function apiBaseUrl() { return process.env.PAPERBOT_API_BASE_URL || "http://127.0.0.1:8000" } +import { withBackendAuth } from "../../_utils/auth-headers" + export async function POST(req: Request) { const body = await req.text() const upstream = await fetch(`${apiBaseUrl()}/api/studio/chat`, { method: "POST", - headers: { + headers: await withBackendAuth(req, { "Content-Type": req.headers.get("content-type") || "application/json", - }, + }), body, }) diff --git a/web/src/app/dashboard/page.tsx b/web/src/app/dashboard/page.tsx index 56c72169..198eceea 100644 --- a/web/src/app/dashboard/page.tsx +++ b/web/src/app/dashboard/page.tsx @@ -10,6 +10,7 @@ import { type LucideIcon, TrendingUp, } from "lucide-react" +import { auth } from "@/auth" import DashboardReadingQueuePanel from "@/components/dashboard/DashboardReadingQueuePanel" import { fetchDeadlineRadar, fetchPapers } from "@/lib/api" @@ -484,6 +485,8 @@ function DeadlineCard({ item }: { item: DeadlineRadarItem }) { export default async function DashboardPage() { noStore() + const session = await auth() + const accessToken = session?.accessToken const [ tracksResult, @@ -493,12 +496,12 @@ export default async function DashboardPage() { latestBriefResult, deadlinesResult, ] = await Promise.allSettled([ - fetchDashboardTracks("default"), - fetchDashboardReadingQueue("default", 6), - fetchIntelligenceFeed("default", 6, { sortBy: "delta", sortOrder: "desc" }), - fetchPapers(), + fetchDashboardTracks(accessToken), + fetchDashboardReadingQueue(accessToken, 6), + fetchIntelligenceFeed(accessToken, 6, { sortBy: "delta", sortOrder: "desc" }), + fetchPapers(accessToken), fetchLatestDashboardBrief(), - fetchDeadlineRadar("default"), + fetchDeadlineRadar(accessToken), ]) const tracks = tracksResult.status === "fulfilled" ? tracksResult.value : [] @@ -513,7 +516,7 @@ export default async function DashboardPage() { const orderedTracks = [...tracks].sort((left, right) => Number(Boolean(right.is_active)) - Number(Boolean(left.is_active))) const activeTrack = orderedTracks[0] || null const activeTrackFeed = activeTrack - ? await fetchDashboardTrackFeed(activeTrack.id, "default", 4).catch(() => ({ items: [], total: 0 })) + ? await fetchDashboardTrackFeed(activeTrack.id, accessToken, 4).catch(() => ({ items: [], total: 0 })) : { items: [], total: 0 } const recommendationCards: DashboardRecommendationCardData[] = latestBrief?.recommendations.length ? latestBrief.recommendations.map((item) => ({ @@ -550,7 +553,7 @@ export default async function DashboardPage() { const spotlightSeeds = orderedTracks.slice(0, 3) const spotlightItems: TrackSpotlightItem[] = await Promise.all( spotlightSeeds.map(async (track) => { - const feed = await fetchDashboardTrackFeed(track.id, "default", 1).catch(() => ({ items: [], total: 0 })) + const feed = await fetchDashboardTrackFeed(track.id, accessToken, 1).catch(() => ({ items: [], total: 0 })) return { id: track.id, name: track.name, diff --git a/web/src/app/forgot-password/page.tsx b/web/src/app/forgot-password/page.tsx new file mode 100644 index 00000000..47daea77 --- /dev/null +++ b/web/src/app/forgot-password/page.tsx @@ -0,0 +1,115 @@ +"use client" + +import { useState } from "react" +import Link from "next/link" +import { BookOpen, Loader2 } from "lucide-react" +import { Button } from "@/components/ui/button" +import { Input } from "@/components/ui/input" +import { Label } from "@/components/ui/label" + +export default function ForgotPasswordPage() { + const [email, setEmail] = useState("") + const [loading, setLoading] = useState(false) + const [submitted, setSubmitted] = useState(false) + const [error, setError] = useState(null) + + const onSubmit = async (e: React.FormEvent) => { + e.preventDefault() + setError(null) + setLoading(true) + try { + await fetch("/api/auth/forgot-password", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ email }), + }) + // Always show success to avoid email enumeration + setSubmitted(true) + } catch { + setError("Something went wrong. Please try again.") + } finally { + setLoading(false) + } + } + + return ( +
+
+
+ + PaperBot +
+ + {submitted ? ( +
+
+

Check your inbox

+

+ If {email} is registered, + you'll receive a reset link shortly. +

+
+

+ Didn't get it? Check your spam folder or{" "} + + . +

+ + Back to sign in + +
+ ) : ( + <> +
+

Forgot your password?

+

+ Enter your email and we'll send you a reset link. +

+
+ +
+
+ + setEmail(e.target.value)} + disabled={loading} + required + /> +
+ + {error &&

{error}

} + + +
+ +

+ Remember your password?{" "} + + Sign in + +

+ + )} +
+
+ ) +} diff --git a/web/src/app/layout.tsx b/web/src/app/layout.tsx index 2aa01bbd..4b53a06a 100644 --- a/web/src/app/layout.tsx +++ b/web/src/app/layout.tsx @@ -2,6 +2,7 @@ import type { Metadata } from "next"; import { Geist, Geist_Mono } from "next/font/google"; import "./globals.css"; import { LayoutShell } from "@/components/layout/LayoutShell"; +import { SessionProvider } from "next-auth/react"; const geistSans = Geist({ variable: "--font-geist-sans", @@ -28,7 +29,9 @@ export default function RootLayout({ - {children} + + {children} + ); diff --git a/web/src/app/login/page.tsx b/web/src/app/login/page.tsx new file mode 100644 index 00000000..0154251d --- /dev/null +++ b/web/src/app/login/page.tsx @@ -0,0 +1,210 @@ +"use client" + +import { useState } from "react" +import { signIn } from "next-auth/react" +import { useRouter, useSearchParams } from "next/navigation" +import Link from "next/link" +import { Loader2, Eye, EyeOff, BookOpen } from "lucide-react" +import { Button } from "@/components/ui/button" +import { Input } from "@/components/ui/input" +import { Label } from "@/components/ui/label" +import { Separator } from "@/components/ui/separator" + +export default function LoginPage() { + const [email, setEmail] = useState("") + const [password, setPassword] = useState("") + const [showPassword, setShowPassword] = useState(false) + const [error, setError] = useState(null) + const [loading, setLoading] = useState(false) + const [githubLoading, setGithubLoading] = useState(false) + const router = useRouter() + const searchParams = useSearchParams() + const callbackUrl = searchParams.get("callbackUrl") || "/dashboard" + + const onSubmit = async (e: React.FormEvent) => { + e.preventDefault() + setError(null) + setLoading(true) + + // Pre-check to surface specific backend errors (e.g. deleted account) + try { + const check = await fetch("/api/auth/login-check", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ email, password }), + }) + if (!check.ok) { + const data = await check.json().catch(() => ({})) as { detail?: string } + setError(data.detail || "Invalid email or password.") + setLoading(false) + return + } + } catch { + // Network error, fall through to signIn which will also fail + } + + const res = await signIn("credentials", { email, password, redirect: false }) + setLoading(false) + if (res?.error) { + setError("Invalid email or password.") + return + } + router.push(callbackUrl) + } + + const onGithub = () => { + setGithubLoading(true) + signIn("github", { callbackUrl }) + } + + return ( +
+ {/* Left branding panel */} +
+
+ + PaperBot +
+ +
+

+ “Stay ahead of the literature.
+ Let the research come to you.” +

+
+ + + +
+
+ +

© 2026 PaperBot. All rights reserved.

+
+ + {/* Right form panel */} +
+
+ {/* Mobile logo */} +
+ + PaperBot +
+ + {/* Heading */} +
+

Welcome back

+

+ Sign in to your account to continue +

+
+ + {/* GitHub */} + + +
+ + or + +
+ + {/* Email / password form */} +
+
+ + setEmail(e.target.value)} + disabled={loading} + required + /> +
+ +
+
+ + + Forgot password? + +
+
+ setPassword(e.target.value)} + disabled={loading} + required + className="pr-9" + /> + +
+
+ + {error && ( +

{error}

+ )} + + +
+ +

+ Don't have an account?{" "} + + Create one + +

+
+
+
+ ) +} + +function GitHubIcon() { + return ( + + ) +} + +function BulletPoint({ text }: { text: string }) { + return ( +
+ + {text} +
+ ) +} diff --git a/web/src/app/papers/[id]/page.tsx b/web/src/app/papers/[id]/page.tsx index 1a98dae7..eff6ff6c 100644 --- a/web/src/app/papers/[id]/page.tsx +++ b/web/src/app/papers/[id]/page.tsx @@ -1,4 +1,5 @@ import { fetchPaperDetails } from "@/lib/api" +import { auth } from "@/auth" import { Badge } from "@/components/ui/badge" import { Button } from "@/components/ui/button" import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs" @@ -14,7 +15,9 @@ import { VelocityChart } from "@/components/paper/VelocityChart" export default async function PaperPage({ params }: { params: Promise<{ id: string }> }) { const { id } = await params - const paper = await fetchPaperDetails(id) + const session = await auth() + const accessToken = (session as any)?.accessToken as string | undefined + const paper = await fetchPaperDetails(id, accessToken) return (
diff --git a/web/src/app/papers/page.tsx b/web/src/app/papers/page.tsx index 0609cee9..cf241c33 100644 --- a/web/src/app/papers/page.tsx +++ b/web/src/app/papers/page.tsx @@ -1,6 +1,14 @@ +import { redirect } from "next/navigation" + +import { auth } from "@/auth" import SavedPapersList from "@/components/research/SavedPapersList" -export default function PapersPage() { +export default async function PapersPage() { + const session = await auth() + if (!session) { + redirect("/login?callbackUrl=/papers") + } + return (
diff --git a/web/src/app/register/page.tsx b/web/src/app/register/page.tsx new file mode 100644 index 00000000..fa5378a0 --- /dev/null +++ b/web/src/app/register/page.tsx @@ -0,0 +1,316 @@ +"use client" + +import { useState, useMemo } from "react" +import { signIn } from "next-auth/react" +import { useRouter } from "next/navigation" +import Link from "next/link" +import { Loader2, Eye, EyeOff, BookOpen } from "lucide-react" +import { Button } from "@/components/ui/button" +import { Input } from "@/components/ui/input" +import { Label } from "@/components/ui/label" +import { Separator } from "@/components/ui/separator" + +// ─── Password strength ──────────────────────────────────────────────────────── + +type Strength = { score: 0 | 1 | 2 | 3 | 4; label: string; color: string } + +function getStrength(pw: string): Strength { + if (!pw) return { score: 0, label: "", color: "" } + let score = 0 + if (pw.length >= 8) score++ + if (pw.length >= 12) score++ + if (/[A-Z]/.test(pw) && /[a-z]/.test(pw)) score++ + if (/\d/.test(pw) && /[^A-Za-z0-9]/.test(pw)) score++ + + const levels: Strength[] = [ + { score: 0, label: "", color: "" }, + { score: 1, label: "Weak", color: "bg-red-500" }, + { score: 2, label: "Fair", color: "bg-orange-400" }, + { score: 3, label: "Good", color: "bg-yellow-400" }, + { score: 4, label: "Strong", color: "bg-green-500" }, + ] + return levels[score as 0 | 1 | 2 | 3 | 4] +} + +// ─── Page ───────────────────────────────────────────────────────────────────── + +export default function RegisterPage() { + const [displayName, setDisplayName] = useState("") + const [email, setEmail] = useState("") + const [password, setPassword] = useState("") + const [confirm, setConfirm] = useState("") + const [showPassword, setShowPassword] = useState(false) + const [showConfirm, setShowConfirm] = useState(false) + const [error, setError] = useState(null) + const [loading, setLoading] = useState(false) + const [githubLoading, setGithubLoading] = useState(false) + const router = useRouter() + + const strength = useMemo(() => getStrength(password), [password]) + const mismatch = confirm.length > 0 && confirm !== password + + const onSubmit = async (e: React.FormEvent) => { + e.preventDefault() + if (password !== confirm) { + setError("Passwords do not match.") + return + } + setError(null) + setLoading(true) + + const res = await fetch("/api/auth/register", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ email, password, display_name: displayName || undefined }), + }) + + if (!res.ok) { + const body = await res.json().catch(() => null) + setError(body?.detail ?? "Registration failed. Please try again.") + setLoading(false) + return + } + + const r = await signIn("credentials", { email, password, redirect: false }) + setLoading(false) + if (r?.error) { + setError("Account created. Please sign in.") + router.push("/login") + } else { + router.push("/dashboard") + } + } + + const onGithub = () => { + setGithubLoading(true) + signIn("github", { callbackUrl: "/dashboard" }) + } + + return ( +
+ {/* Left branding panel */} +
+
+ + PaperBot +
+ +
+

+ “Your personal AI research assistant,
+ built for serious readers.” +

+
+ + + +
+
+ +

© 2026 PaperBot. All rights reserved.

+
+ + {/* Right form panel */} +
+
+ {/* Mobile logo */} +
+ + PaperBot +
+ + {/* Heading */} +
+

Create an account

+

+ Get started for free in under a minute +

+
+ + {/* GitHub */} + + +
+ + or + +
+ + {/* Form */} +
+
+ + setDisplayName(e.target.value)} + disabled={loading} + /> +
+ +
+ + setEmail(e.target.value)} + disabled={loading} + required + /> +
+ + {/* Password + strength meter */} +
+ +
+ setPassword(e.target.value)} + disabled={loading} + required + minLength={8} + className="pr-9" + /> + +
+ + {/* Strength bar — only visible once the user starts typing */} + {password.length > 0 && ( +
+
+ {([1, 2, 3, 4] as const).map(n => ( +
= n ? strength.color : "bg-muted" + }`} + /> + ))} +
+ {strength.label && ( +

+ Strength:{" "} + + {strength.label} + + {strength.score < 3 && ( + + — try adding uppercase, numbers, or symbols + + )} +

+ )} +
+ )} +
+ + {/* Confirm password */} +
+ +
+ setConfirm(e.target.value)} + disabled={loading} + required + className={`pr-9 ${mismatch ? "border-destructive focus-visible:ring-destructive/30" : ""}`} + /> + +
+ {mismatch && ( +

Passwords do not match.

+ )} +
+ + {error && ( +

{error}

+ )} + + + + +

+ Already have an account?{" "} + + Sign in + +

+
+
+
+ ) +} + +function GitHubIcon() { + return ( + + ) +} + +function BulletPoint({ text }: { text: string }) { + return ( +
+ + {text} +
+ ) +} diff --git a/web/src/app/research/page.tsx b/web/src/app/research/page.tsx index 94ab648b..8ce28c06 100644 --- a/web/src/app/research/page.tsx +++ b/web/src/app/research/page.tsx @@ -1,8 +1,15 @@ import { Suspense } from "react" +import { redirect } from "next/navigation" +import { auth } from "@/auth" import ResearchPageNew from "@/components/research/ResearchPageNew" -export default function ResearchPage() { +export default async function ResearchPage() { + const session = await auth() + if (!session) { + redirect("/login?callbackUrl=/research") + } + return (
Loading research workspace...
}> diff --git a/web/src/app/reset-password/page.tsx b/web/src/app/reset-password/page.tsx new file mode 100644 index 00000000..22ec16ea --- /dev/null +++ b/web/src/app/reset-password/page.tsx @@ -0,0 +1,190 @@ +"use client" + +import { useState, useMemo, Suspense } from "react" +import { useSearchParams, useRouter } from "next/navigation" +import Link from "next/link" +import { BookOpen, Eye, EyeOff, Loader2 } from "lucide-react" +import { Button } from "@/components/ui/button" +import { Input } from "@/components/ui/input" +import { Label } from "@/components/ui/label" + +type Strength = { score: 0 | 1 | 2 | 3 | 4; label: string; color: string } + +function getStrength(pw: string): Strength { + if (!pw) return { score: 0, label: "", color: "" } + let score = 0 + if (pw.length >= 8) score++ + if (pw.length >= 12) score++ + if (/[A-Z]/.test(pw) && /[a-z]/.test(pw)) score++ + if (/\d/.test(pw) && /[^A-Za-z0-9]/.test(pw)) score++ + const levels: Strength[] = [ + { score: 0, label: "", color: "" }, + { score: 1, label: "Weak", color: "bg-red-500" }, + { score: 2, label: "Fair", color: "bg-orange-400" }, + { score: 3, label: "Good", color: "bg-yellow-400" }, + { score: 4, label: "Strong", color: "bg-green-500" }, + ] + return levels[score as 0 | 1 | 2 | 3 | 4] +} + +function ResetPasswordForm() { + const searchParams = useSearchParams() + const token = searchParams.get("token") ?? "" + const router = useRouter() + + const [password, setPassword] = useState("") + const [confirm, setConfirm] = useState("") + const [showPassword, setShowPassword] = useState(false) + const [loading, setLoading] = useState(false) + const [error, setError] = useState(null) + const [done, setDone] = useState(false) + + const strength = useMemo(() => getStrength(password), [password]) + const mismatch = confirm.length > 0 && confirm !== password + + if (!token) { + return ( +
+

Invalid or missing reset token.

+ + Request a new link + +
+ ) + } + + const onSubmit = async (e: React.FormEvent) => { + e.preventDefault() + if (password !== confirm) { setError("Passwords do not match."); return } + setError(null) + setLoading(true) + try { + const res = await fetch("/api/auth/reset-password", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ token, new_password: password }), + }) + if (!res.ok) { + const body = await res.json().catch(() => null) + setError(body?.detail ?? "Failed to reset password. The link may have expired.") + return + } + setDone(true) + setTimeout(() => router.push("/login"), 2500) + } catch { + setError("Something went wrong. Please try again.") + } finally { + setLoading(false) + } + } + + if (done) { + return ( +
+

Password updated!

+

Redirecting you to sign in…

+
+ ) + } + + return ( + <> +
+

Set a new password

+

Choose a strong password for your account.

+
+ +
+
+ +
+ setPassword(e.target.value)} + disabled={loading} + required + minLength={8} + className="pr-9" + /> + +
+ {password.length > 0 && ( +
+
+ {([1, 2, 3, 4] as const).map(n => ( +
= n ? strength.color : "bg-muted" + }`} + /> + ))} +
+ {strength.label && ( +

+ Strength:{" "} + {strength.label} +

+ )} +
+ )} +
+ +
+ + setConfirm(e.target.value)} + disabled={loading} + required + className={mismatch ? "border-destructive focus-visible:ring-destructive/30" : ""} + /> + {mismatch &&

Passwords do not match.

} +
+ + {error &&

{error}

} + + + + + ) +} + +export default function ResetPasswordPage() { + return ( +
+
+
+ + PaperBot +
+ Loading…

}> + +
+
+
+ ) +} diff --git a/web/src/app/settings/page.tsx b/web/src/app/settings/page.tsx index 3ec22d00..107f4a45 100644 --- a/web/src/app/settings/page.tsx +++ b/web/src/app/settings/page.tsx @@ -1,12 +1,26 @@ "use client" import { useEffect, useMemo, useState } from "react" -import { CheckCircle2, KeyRound, Loader2, Plus, Trash2, Wrench, Pencil } from "lucide-react" +import { + CheckCircle2, + Eye, + EyeOff, + KeyRound, + Loader2, + Pencil, + Plus, + Trash2, + Wrench, +} from "lucide-react" +import { useSession, signOut } from "next-auth/react" import { Badge } from "@/components/ui/badge" import { Button } from "@/components/ui/button" -import { Card, CardContent } from "@/components/ui/card" +import { Card, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card" import { Input } from "@/components/ui/input" +import { Label } from "@/components/ui/label" +import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs" +import { fetchJson, getErrorMessage } from "@/lib/fetch" import { Dialog, DialogContent, @@ -17,7 +31,8 @@ import { } from "@/components/ui/dialog" import { EmbeddingSettingsPanel } from "@/components/settings/EmbeddingSettingsPanel" -import { ScholarSubscriptionsPanel } from "@/components/settings/ScholarSubscriptionsPanel" + +// ─── Types ──────────────────────────────────────────────────────────────────── type ModelEndpoint = { id: number @@ -57,6 +72,8 @@ type Preset = { task_types: string[] } +// ─── Constants ──────────────────────────────────────────────────────────────── + const QUICK_PRESETS: Preset[] = [ { label: "OpenAI", @@ -108,6 +125,8 @@ const EMPTY_FORM: FormState = { is_default: false, } +// ─── Helpers ────────────────────────────────────────────────────────────────── + function toPayload(form: FormState) { return { name: form.name.trim(), @@ -134,7 +153,299 @@ function statusDot(item: ModelEndpoint) { return "bg-gray-400" } -export default function SettingsPage() { +// ─── Account Section ────────────────────────────────────────────────────────── + +function AccountSection() { + const { data: session, update: updateSession } = useSession() + + const isOAuth = session?.provider === "github" + + const [displayName, setDisplayName] = useState("") + const [nameLoading, setNameLoading] = useState(false) + const [nameMsg, setNameMsg] = useState<{ type: "ok" | "err"; text: string } | null>(null) + + const [currentPw, setCurrentPw] = useState("") + const [newPw, setNewPw] = useState("") + const [confirmPw, setConfirmPw] = useState("") + const [showCurrentPw, setShowCurrentPw] = useState(false) + const [showNewPw, setShowNewPw] = useState(false) + const [pwLoading, setPwLoading] = useState(false) + const [pwMsg, setPwMsg] = useState<{ type: "ok" | "err"; text: string } | null>(null) + + const [deleteDialogOpen, setDeleteDialogOpen] = useState(false) + const [deleteConfirmText, setDeleteConfirmText] = useState("") + const [deleteLoading, setDeleteLoading] = useState(false) + const [deleteErr, setDeleteErr] = useState(null) + + useEffect(() => { + setDisplayName(session?.user?.name || "") + }, [session]) + + async function saveName() { + setNameLoading(true) + setNameMsg(null) + try { + const data = await fetchJson<{ display_name?: string }>("/api/auth/me", { + method: "PATCH", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ display_name: displayName.trim() || null }), + }) + await updateSession({ name: data.display_name ?? displayName.trim() }) + setNameMsg({ type: "ok", text: "Display name updated." }) + } catch (e) { + setNameMsg({ type: "err", text: getErrorMessage(e) }) + } finally { + setNameLoading(false) + } + } + + async function changePassword() { + if (newPw !== confirmPw) { + setPwMsg({ type: "err", text: "New passwords do not match." }) + return + } + setPwLoading(true) + setPwMsg(null) + try { + await fetchJson("/api/auth/me/change-password", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ current_password: currentPw, new_password: newPw }), + }) + setCurrentPw("") + setNewPw("") + setConfirmPw("") + setPwMsg({ type: "ok", text: "Password updated successfully." }) + } catch (e) { + setPwMsg({ type: "err", text: getErrorMessage(e) }) + } finally { + setPwLoading(false) + } + } + + async function deleteAccount() { + setDeleteLoading(true) + setDeleteErr(null) + try { + const res = await fetch("/api/auth/me", { method: "DELETE" }) + if (!res.ok) { + const data = await res.json().catch(() => null) + throw new Error(data?.detail || `${res.status}`) + } + await signOut({ callbackUrl: "/login" }) + } catch (e) { + setDeleteErr(getErrorMessage(e)) + setDeleteLoading(false) + } + } + + const email = session?.user?.email ?? undefined + + return ( +
+ {/* Profile */} + + + Profile + Update your display name shown across the app. + + + {email && ( +
+ + +
+ )} +
+ +
+ setDisplayName(e.target.value)} + placeholder="Your name" + disabled={nameLoading} + /> + +
+ {nameMsg && ( +

+ {nameMsg.text} +

+ )} +
+
+
+ + {/* Security */} + + + Security + + {isOAuth + ? "Your account is managed via GitHub. Password sign-in is not available." + : "Change your password. You'll stay signed in on this device."} + + + + {isOAuth ? ( +

Signed in with GitHub — no password to manage.

+ ) : ( +
+
+ +
+ setCurrentPw(e.target.value)} + disabled={pwLoading} + className="pr-9" + /> + +
+
+ +
+ +
+ setNewPw(e.target.value)} + disabled={pwLoading} + minLength={8} + placeholder="Min. 8 characters" + className="pr-9" + /> + +
+
+ +
+ + setConfirmPw(e.target.value)} + disabled={pwLoading} + className={confirmPw && confirmPw !== newPw ? "border-destructive" : ""} + /> + {confirmPw && confirmPw !== newPw && ( +

Passwords do not match.

+ )} +
+ + {pwMsg && ( +

+ {pwMsg.text} +

+ )} + + +
+ )} +
+
+ + {/* Danger Zone */} + + + Danger zone + These actions are irreversible. Please proceed with caution. + + +
+
+

Sign out

+

End your current session.

+
+ +
+ +
+
+

Delete account

+

+ Permanently deactivate your account. Your data will no longer be accessible. +

+
+ +
+
+
+ + {/* Delete confirmation dialog */} + + + + Delete your account? + + This will permanently deactivate your account. This action cannot be undone. + Type delete my account to confirm. + + + setDeleteConfirmText(e.target.value)} + placeholder="delete my account" + disabled={deleteLoading} + /> + {deleteErr &&

{deleteErr}

} + + + + +
+
+
+ ) +} + +// ─── Model Providers Section ────────────────────────────────────────────────── + +function ModelProvidersSection() { const [items, setItems] = useState([]) const [loading, setLoading] = useState(false) const [saving, setSaving] = useState(false) @@ -155,7 +466,7 @@ export default function SettingsPage() { const payload = await res.json() setItems(payload.items || []) } catch (e) { - setError(e instanceof Error ? e.message : String(e)) + setError(getErrorMessage(e)) setItems([]) } finally { setLoading(false) @@ -219,7 +530,7 @@ export default function SettingsPage() { setDialogOpen(false) setMessage(editing ? "Provider updated." : "Provider created.") } catch (e) { - setError(e instanceof Error ? e.message : String(e)) + setError(getErrorMessage(e)) } finally { setSaving(false) } @@ -235,7 +546,7 @@ export default function SettingsPage() { await load() setMessage("Provider removed.") } catch (e) { - setError(e instanceof Error ? e.message : String(e)) + setError(getErrorMessage(e)) } } @@ -248,7 +559,7 @@ export default function SettingsPage() { await load() setMessage("Provider activated.") } catch (e) { - setError(e instanceof Error ? e.message : String(e)) + setError(getErrorMessage(e)) } } @@ -266,25 +577,17 @@ export default function SettingsPage() { if (!res.ok) throw new Error(String(payload?.detail || `${res.status}`)) setMessage(payload?.message || "Connection test passed.") } catch (e) { - setError(e instanceof Error ? e.message : String(e)) + setError(getErrorMessage(e)) } finally { setTestingId(null) } } return ( -
-

Settings

- -
-

Model Providers

-

Configure LLM providers for paper analysis

-
- + <> {error &&

{error}

} {message &&

{message}

} - {/* Provider Cards */}
{loading ? (

Loading providers...

@@ -303,9 +606,7 @@ export default function SettingsPage() {
{item.name} - {item.is_default && ( - Default - )} + {item.is_default && Default}

{(item.models || []).join(", ") || "no models"} · {item.base_url || "(default URL)"} @@ -329,13 +630,7 @@ export default function SettingsPage() {

-
- {/* Add Provider + Quick Presets */}
-

Quick Presets

@@ -379,9 +672,7 @@ export default function SettingsPage() {
-
- -
+ {/* Add/Edit Dialog */} @@ -472,6 +763,36 @@ export default function SettingsPage() { + + ) +} + +// ─── Page ───────────────────────────────────────────────────────────────────── + +export default function SettingsPage() { + return ( +
+

Settings

+ + + + Account + Model Providers + + + + + + + +
+
+

Configure LLM providers for paper analysis.

+
+ +
+
+
) } diff --git a/web/src/auth.ts b/web/src/auth.ts new file mode 100644 index 00000000..aed2aad1 --- /dev/null +++ b/web/src/auth.ts @@ -0,0 +1,93 @@ +import NextAuth from "next-auth" +import GitHub from "next-auth/providers/github" +import Credentials from "next-auth/providers/credentials" + +function backendBaseUrl() { + return process.env.BACKEND_BASE_URL || process.env.PAPERBOT_API_BASE_URL || "http://127.0.0.1:8000" +} + +export const { handlers, signIn, signOut, auth } = NextAuth({ + providers: [ + GitHub, + Credentials({ + credentials: { + email: { label: "Email", type: "email" }, + password: { label: "Password", type: "password" }, + }, + async authorize(credentials) { + try { + const res = await fetch(`${backendBaseUrl()}/api/auth/login`, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ email: credentials?.email, password: credentials?.password }), + }) + if (!res.ok) return null + const data = await res.json() as { access_token: string; user_id: number; display_name?: string } + return { id: String(data.user_id), name: data.display_name || "", accessToken: data.access_token, userId: data.user_id } as any + } catch { + return null + } + }, + }), + ], + callbacks: { + async jwt({ token, user, account, profile, trigger, session: sessionData }) { + // Handle updateSession() calls — persist updated fields into the token + if (trigger === "update") { + if (sessionData?.name !== undefined) token.name = sessionData.name + } + + // Credentials: copy backend JWT from user + if (user && (user as any).accessToken) { + token.accessToken = (user as any).accessToken + token.userId = (user as any).userId + if (user.name) token.name = user.name + } + + // Persist provider on first sign-in + if (account?.provider) { + token.provider = account.provider + } + + // GitHub OAuth: exchange for backend JWT + if (account?.provider === "github" && account.access_token && profile) { + const res = await fetch(`${backendBaseUrl()}/api/auth/github/exchange`, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ + github_id: String((profile as any).id), + login: (profile as any).login, + name: (profile as any).name, + avatar_url: (profile as any).avatar_url, + email: (profile as any).email, + access_token: account.access_token, + }), + }) + if (res.ok) { + const data = await res.json() as { access_token: string; user_id: number; display_name?: string } + token.accessToken = data.access_token + token.userId = data.user_id + if (data.display_name) token.name = data.display_name + } else { + console.error("[auth] github/exchange failed:", res.status, await res.text().catch(() => "")) + // Ensure we do not keep a half-authenticated session without backend JWT + delete token.accessToken + delete token.userId + } + } + return token + }, + async session({ session, token }) { + ;(session as any).accessToken = token.accessToken + ;(session as any).userId = token.userId + ;(session as any).provider = token.provider + if (token.name !== undefined) session.user.name = token.name as string + return session + }, + }, + pages: { + signIn: "/login", + error: "/login", + }, +}) + diff --git a/web/src/components/dashboard/ActivityFeed.tsx b/web/src/components/dashboard/ActivityFeed.tsx index ddd0c781..a7454c06 100644 --- a/web/src/components/dashboard/ActivityFeed.tsx +++ b/web/src/components/dashboard/ActivityFeed.tsx @@ -1,5 +1,6 @@ import type { Activity } from "@/lib/types" import { fetchActivities } from "@/lib/api" +import { auth } from "@/auth" import { NewPaperCard } from "./feed/NewPaperCard" import { MilestoneCard } from "./feed/MilestoneCard" @@ -12,7 +13,11 @@ interface ActivityFeedProps { } export async function ActivityFeed({ activities, maxItems = 6, showTitle = true }: ActivityFeedProps) { - const items = activities ?? (await fetchActivities()) + const items = activities ?? (await (async () => { + const session = await auth() + const accessToken = (session as any)?.accessToken as string | undefined + return fetchActivities(accessToken) + })()) const visible = items.slice(0, maxItems) return ( diff --git a/web/src/components/dashboard/DashboardReadingQueuePanel.tsx b/web/src/components/dashboard/DashboardReadingQueuePanel.tsx index abbabe8d..ba01022a 100644 --- a/web/src/components/dashboard/DashboardReadingQueuePanel.tsx +++ b/web/src/components/dashboard/DashboardReadingQueuePanel.tsx @@ -108,8 +108,8 @@ function QueueCard({ setError(null) try { const [detailRes, cardRes] = await Promise.all([ - fetch(`/api/research/papers/${encodeURIComponent(detailRef)}?user_id=default`, { cache: "no-store" }), - fetch(`/api/research/papers/${encodeURIComponent(detailRef)}/card?user_id=default`, { cache: "no-store" }), + fetch(`/api/research/papers/${encodeURIComponent(detailRef)}`, { cache: "no-store" }), + fetch(`/api/research/papers/${encodeURIComponent(detailRef)}/card`, { cache: "no-store" }), ]) const detailPayload = detailRes.ok ? await detailRes.json() : null @@ -141,7 +141,6 @@ function QueueCard({ method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ - user_id: "default", track_id: activeTrackId, }), }) @@ -154,7 +153,6 @@ function QueueCard({ method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ - user_id: "default", track_id: activeTrackId, paper_id: item.paperRef, action: "save", @@ -189,7 +187,6 @@ function QueueCard({ method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ - user_id: "default", status: "archived", mark_saved: true, }), diff --git a/web/src/components/dashboard/TrackSpotlightSection.tsx b/web/src/components/dashboard/TrackSpotlightSection.tsx index 7d0983ce..a6726519 100644 --- a/web/src/components/dashboard/TrackSpotlightSection.tsx +++ b/web/src/components/dashboard/TrackSpotlightSection.tsx @@ -11,7 +11,6 @@ interface TrackSpotlightSectionProps { initialFeedItems: TrackFeedItem[] initialFeedTotal: number initialAnchors: AnchorPreviewItem[] - userId?: string } export function TrackSpotlightSection({ @@ -19,8 +18,7 @@ export function TrackSpotlightSection({ initialActiveTrack, initialFeedItems, initialFeedTotal, - initialAnchors, - userId = "default", + initialAnchors = "default", }: TrackSpotlightSectionProps) { const [tracks] = useState(initialTracks) const [activeTrack, setActiveTrack] = useState(initialActiveTrack) @@ -40,8 +38,8 @@ export function TrackSpotlightSection({ try { // Fetch feed and anchors for the new track const [feedRes, anchorsRes] = await Promise.all([ - fetch(`/api/research/tracks/${trackId}/feed?user_id=${encodeURIComponent(userId)}&limit=6`), - fetch(`/api/research/tracks/${trackId}/anchors/discover?user_id=${encodeURIComponent(userId)}&limit=4`), + fetch(`/api/research/tracks/${trackId}/feed?limit=6`), + fetch(`/api/research/tracks/${trackId}/anchors/discover?limit=4`), ]) if (feedRes.ok) { @@ -61,7 +59,7 @@ export function TrackSpotlightSection({ } // Optionally activate the track on the backend - await fetch(`/api/research/tracks/${trackId}/activate?user_id=${encodeURIComponent(userId)}`, { + await fetch(`/api/research/tracks/${trackId}/activate`, { method: "POST", headers: { "Content-Type": "application/json" }, body: "{}", @@ -75,7 +73,7 @@ export function TrackSpotlightSection({ setIsLoading(false) } }, - [tracks, activeTrack?.id, userId] + [tracks, activeTrack?.id] ) return ( diff --git a/web/src/components/layout/Sidebar.tsx b/web/src/components/layout/Sidebar.tsx index adec9bc7..bbb347fb 100644 --- a/web/src/components/layout/Sidebar.tsx +++ b/web/src/components/layout/Sidebar.tsx @@ -2,8 +2,17 @@ import Link from "next/link" import { usePathname } from "next/navigation" +import { useSession, signOut } from "next-auth/react" import { cn } from "@/lib/utils" import { Button } from "@/components/ui/button" +import { + DropdownMenu, + DropdownMenuContent, + DropdownMenuItem, + DropdownMenuLabel, + DropdownMenuSeparator, + DropdownMenuTrigger, +} from "@/components/ui/dropdown-menu" import { LayoutDashboard, Users, @@ -16,6 +25,10 @@ import { PanelLeftClose, PanelLeft, Rocket, + User as UserIcon, + LogOut, + LogIn, + ChevronUp, } from "lucide-react" type SidebarProps = React.HTMLAttributes & { @@ -37,6 +50,14 @@ const routes = [ export function Sidebar({ className, collapsed, onToggle }: SidebarProps) { const pathname = usePathname() const demoUrl = process.env.NEXT_PUBLIC_DEMO_URL + const { data: session } = useSession() + + const isAuthenticated = !!session + const displayName = + session?.user?.name || + session?.user?.email || + String(session?.userId ?? "") || + "Guest" return (
@@ -82,8 +103,69 @@ export function Sidebar({ className, collapsed, onToggle }: SidebarProps) {
- {demoUrl ? ( -
+
+ + + + + + +
+ {displayName} + {isAuthenticated && ( + + {session?.user?.email} + + )} +
+
+ + + + + Settings + + + + {isAuthenticated ? ( + signOut({ callbackUrl: "/login" })} + > + + Sign out + + ) : ( + + + + Sign in + + + )} +
+
+ + {demoUrl ? ( -
- ) : null} + ) : null} +
) } diff --git a/web/src/components/research/DiscoveryGraphWorkspace.tsx b/web/src/components/research/DiscoveryGraphWorkspace.tsx index 3e8635e7..08d5df8a 100644 --- a/web/src/components/research/DiscoveryGraphWorkspace.tsx +++ b/web/src/components/research/DiscoveryGraphWorkspace.tsx @@ -94,7 +94,6 @@ function toValidYear(value: string): number | undefined { } interface DiscoveryGraphWorkspaceProps { - userId: string trackId: number | null onSavePaper: (paper: SavePaperPayload) => Promise initialSeedType?: SeedType @@ -103,8 +102,7 @@ interface DiscoveryGraphWorkspaceProps { } export default function DiscoveryGraphWorkspace({ - userId, - trackId, + trackId, onSavePaper, initialSeedType = "doi", initialSeedId = "", @@ -225,7 +223,6 @@ export default function DiscoveryGraphWorkspace({ method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ - user_id: userId, track_id: trackId ?? undefined, seed_type: seedType, seed_id: seedId.trim(), diff --git a/web/src/components/research/FeedTab.tsx b/web/src/components/research/FeedTab.tsx index 0dad48af..5fb42091 100644 --- a/web/src/components/research/FeedTab.tsx +++ b/web/src/components/research/FeedTab.tsx @@ -40,7 +40,6 @@ type FeedResponse = { } interface FeedTabProps { - userId: string trackId: number | null onFeedbackAction?: ( paperId: string, @@ -72,7 +71,7 @@ function toPaper(item: FeedItem): Paper { } } -export function FeedTab({ userId, trackId, onFeedbackAction }: FeedTabProps) { +export function FeedTab({ trackId, onFeedbackAction }: FeedTabProps) { const [items, setItems] = useState([]) const [loading, setLoading] = useState(false) const [error, setError] = useState(null) @@ -88,11 +87,10 @@ export function FeedTab({ userId, trackId, onFeedbackAction }: FeedTabProps) { setError(null) try { const qs = new URLSearchParams({ - user_id: userId, limit: "20", offset: "0", }) - const res = await fetch(`/api/research/tracks/${trackId}/feed?${qs.toString()}`) + const res = await fetch(`/api/research/tracks/${trackId}/feed`) if (!res.ok) { throw new Error(`${res.status} ${res.statusText}`) } @@ -109,7 +107,7 @@ export function FeedTab({ userId, trackId, onFeedbackAction }: FeedTabProps) { useEffect(() => { load().catch(() => {}) // eslint-disable-next-line react-hooks/exhaustive-deps - }, [userId, trackId]) + }, [trackId]) if (!trackId) { return
Select a track to view feed.
diff --git a/web/src/components/research/MemoryTab.tsx b/web/src/components/research/MemoryTab.tsx index 82ab6875..abc6db26 100644 --- a/web/src/components/research/MemoryTab.tsx +++ b/web/src/components/research/MemoryTab.tsx @@ -22,11 +22,10 @@ type MemoryResponse = { } interface MemoryTabProps { - userId: string trackId: number | null } -export function MemoryTab({ userId, trackId }: MemoryTabProps) { +export function MemoryTab({ trackId }: MemoryTabProps) { const [items, setItems] = useState([]) const [loading, setLoading] = useState(false) const [error, setError] = useState(null) @@ -40,11 +39,10 @@ export function MemoryTab({ userId, trackId }: MemoryTabProps) { setError(null) try { const qs = new URLSearchParams({ - user_id: userId, track_id: String(trackId), limit: "100", }) - const res = await fetch(`/api/research/memory/inbox?${qs.toString()}`) + const res = await fetch(`/api/research/memory/inbox`) if (!res.ok) { throw new Error(`${res.status} ${res.statusText}`) } @@ -61,7 +59,7 @@ export function MemoryTab({ userId, trackId }: MemoryTabProps) { useEffect(() => { load().catch(() => {}) // eslint-disable-next-line react-hooks/exhaustive-deps - }, [userId, trackId]) + }, [trackId]) if (!trackId) { return
Select a track to view memory items.
diff --git a/web/src/components/research/ResearchDashboard.tsx b/web/src/components/research/ResearchDashboard.tsx index afc30e07..bd5f91a9 100644 --- a/web/src/components/research/ResearchDashboard.tsx +++ b/web/src/components/research/ResearchDashboard.tsx @@ -1,6 +1,7 @@ "use client" import { useEffect, useMemo, useState } from "react" +import { useSession } from "next-auth/react" import { fetchJson, getErrorMessage } from "@/lib/fetch" import { Badge } from "@/components/ui/badge" @@ -96,7 +97,7 @@ function clampNumber(value: number, min: number, max: number, fallback: number) } export default function ResearchDashboard() { - const [userId, setUserId] = useState("default") + const { data: session } = useSession() const [tracks, setTracks] = useState([]) const [activeTrackId, setActiveTrackId] = useState(null) const [query, setQuery] = useState("") @@ -183,7 +184,7 @@ export default function ResearchDashboard() { } async function refreshTracks(): Promise { - const data = await fetchJson<{ tracks: Track[] }>(`/api/research/tracks?user_id=${encodeURIComponent(userId)}`) + const data = await fetchJson<{ tracks: Track[] }>(`/api/research/tracks`) setTracks(data.tracks || []) const active = data.tracks.find((t) => t.is_active) const activeId = active?.id ?? null @@ -194,17 +195,16 @@ export default function ResearchDashboard() { async function refreshInbox(trackId?: number | null) { const tid = trackId ?? activeTrackId - const qs = new URLSearchParams({ user_id: userId }) if (tid) qs.set("track_id", String(tid)) - const data = await fetchJson<{ items: MemoryItem[] }>(`/api/research/memory/inbox?${qs.toString()}`) + const data = await fetchJson<{ items: MemoryItem[] }>(`/api/research/memory/inbox`) setInbox(data.items || []) setSelectedInboxIds(new Set()) } async function refreshEval() { - const qs = new URLSearchParams({ user_id: userId, days: String(evalDays) }) + const qs = new URLSearchParams({ days: String(evalDays) }) if (activeTrackId) qs.set("track_id", String(activeTrackId)) - const data = await fetchJson<{ summary: EvalSummary }>(`/api/research/evals/summary?${qs.toString()}`) + const data = await fetchJson<{ summary: EvalSummary }>(`/api/research/evals/summary`) setEvalSummary(data.summary) } @@ -219,10 +219,10 @@ export default function ResearchDashboard() { refreshInbox(activeTrackId).catch(() => {}) } // eslint-disable-next-line react-hooks/exhaustive-deps - }, [activeTrackId, userId]) + }, [activeTrackId]) async function activateTrack(trackId: number) { - await fetchJson(`/api/research/tracks/${trackId}/activate?user_id=${encodeURIComponent(userId)}`, { + await fetchJson(`/api/research/tracks/${trackId}/activate`, { method: "POST", body: "{}", headers: { "Content-Type": "application/json" }, @@ -238,7 +238,6 @@ export default function ResearchDashboard() { const suggestion = contextPack?.routing?.suggestion const activateTrackId = activateSuggestion ? suggestion?.track_id : null const body: Record = { - user_id: userId, query, paper_limit: 8, memory_limit: 8, @@ -276,7 +275,6 @@ export default function ResearchDashboard() { await fetchJson(`/api/research/memory/suggest`, { method: "POST", body: JSON.stringify({ - user_id: userId, text: suggestText, scope_type: "track", scope_id: activeTrackId ? String(activeTrackId) : undefined, @@ -303,7 +301,6 @@ export default function ResearchDashboard() { await fetchJson(`/api/research/memory/bulk_moderate`, { method: "POST", body: JSON.stringify({ - user_id: userId, item_ids: ids, status, }), @@ -326,7 +323,6 @@ export default function ResearchDashboard() { await fetchJson(`/api/research/memory/bulk_move`, { method: "POST", body: JSON.stringify({ - user_id: userId, item_ids: ids, scope_type: "track", scope_id: String(targetTrackId), @@ -359,7 +355,6 @@ export default function ResearchDashboard() { await fetchJson(`/api/research/tracks`, { method: "POST", body: JSON.stringify({ - user_id: userId, name, description: newTrackDescription.trim(), keywords, @@ -386,7 +381,6 @@ export default function ResearchDashboard() { try { const contextRunId = contextPack?.context_run_id ?? null const body: Record = { - user_id: userId, track_id: activeTrackId, paper_id: paperId, action, @@ -423,7 +417,7 @@ export default function ResearchDashboard() { setError(null) try { await fetchJson( - `/api/research/tracks/${trackId}/memory/clear?user_id=${encodeURIComponent(userId)}&confirm=true`, + `/api/research/tracks/${trackId}/memory/clear?confirm=true`, { method: "POST", body: "{}", headers: { "Content-Type": "application/json" } }, ) await refreshInbox(trackId) diff --git a/web/src/components/research/ResearchDiscoveryPage.tsx b/web/src/components/research/ResearchDiscoveryPage.tsx index c9f2b29e..c30f9f0c 100644 --- a/web/src/components/research/ResearchDiscoveryPage.tsx +++ b/web/src/components/research/ResearchDiscoveryPage.tsx @@ -2,6 +2,7 @@ import Link from "next/link" import { useEffect, useMemo, useState } from "react" +import { useSession } from "next-auth/react" import { useSearchParams } from "next/navigation" import { ArrowLeft, Compass } from "lucide-react" @@ -15,8 +16,8 @@ import type { Track } from "./TrackSelector" type SeedType = "doi" | "arxiv" | "openalex" | "semantic_scholar" | "author" export default function ResearchDiscoveryPage() { + const { data: session } = useSession() const searchParams = useSearchParams() - const [userId] = useState("default") const [tracks, setTracks] = useState([]) const [activeTrackId, setActiveTrackId] = useState(null) const [loading, setLoading] = useState(false) @@ -59,7 +60,7 @@ export default function ResearchDiscoveryPage() { async function refreshTracks() { const data = await fetchJson<{ tracks: Track[] }>( - `/api/research/tracks?user_id=${encodeURIComponent(userId)}` + `/api/research/tracks` ) setTracks(data.tracks || []) const active = data.tracks.find((track) => track.is_active) @@ -71,7 +72,7 @@ export default function ResearchDiscoveryPage() { setError(null) try { await fetchJson( - `/api/research/tracks/${trackId}/activate?user_id=${encodeURIComponent(userId)}`, + `/api/research/tracks/${trackId}/activate`, { method: "POST", body: "{}", @@ -100,7 +101,6 @@ export default function ResearchDiscoveryPage() { await fetchJson(`/api/research/papers/feedback`, { method: "POST", body: JSON.stringify({ - user_id: userId, track_id: activeTrackId, paper_id: paper.paper_id, action: "save", @@ -175,7 +175,6 @@ export default function ResearchDiscoveryPage() { ([]) @@ -145,7 +146,7 @@ export default function ResearchPageNew() { async function refreshTrackContext(trackId: number): Promise { const data = await fetchJson( - `/api/research/tracks/${trackId}/context?user_id=${encodeURIComponent(userId)}` + `/api/research/tracks/${trackId}/context` ) setTrackContext(data) } @@ -157,7 +158,7 @@ export default function ResearchPageNew() { setTrackContextLoading(true) try { const data = await fetchJson( - `/api/research/tracks/${trackId}/context?user_id=${encodeURIComponent(userId)}` + `/api/research/tracks/${trackId}/context` ) if (!cancelled) { setTrackContext(data) @@ -186,11 +187,11 @@ export default function ResearchPageNew() { return () => { cancelled = true } - }, [activeTrackId, userId]) + }, [activeTrackId]) async function refreshTracks(): Promise { const data = await fetchJson<{ tracks: Track[] }>( - `/api/research/tracks?user_id=${encodeURIComponent(userId)}` + `/api/research/tracks` ) setTracks(data.tracks || []) const active = data.tracks.find((t) => t.is_active) @@ -203,7 +204,7 @@ export default function ResearchPageNew() { setLoading(true) try { await fetchJson( - `/api/research/tracks/${trackId}/activate?user_id=${encodeURIComponent(userId)}`, + `/api/research/tracks/${trackId}/activate`, { method: "POST", body: "{}", @@ -238,7 +239,6 @@ export default function ResearchPageNew() { const parsedYearTo = parseYear(yearTo) const body = { - user_id: userId, query, track_id: activeTrackId ?? undefined, paper_limit: 10, @@ -321,7 +321,6 @@ export default function ResearchPageNew() { await fetchJson(`/api/research/tracks`, { method: "POST", body: JSON.stringify({ - user_id: userId, name, description: data.description, keywords: data.keywords, @@ -368,7 +367,7 @@ export default function ResearchPageNew() { setError(null) setEditError(null) try { - await fetchJson(`/api/research/tracks/${trackId}?user_id=${encodeURIComponent(userId)}`, { + await fetchJson(`/api/research/tracks/${trackId}`, { method: "PATCH", body: JSON.stringify({ name, @@ -407,7 +406,7 @@ export default function ResearchPageNew() { setError(null) try { await fetchJson( - `/api/research/tracks/${trackToClear}/memory/clear?user_id=${encodeURIComponent(userId)}&confirm=true`, + `/api/research/tracks/${trackToClear}/memory/clear?confirm=true`, { method: "POST", body: "{}", @@ -435,7 +434,6 @@ export default function ResearchPageNew() { // Don't set global loading - PaperCard handles its own loading state setError(null) const body: Record = { - user_id: userId, track_id: activeTrackId, paper_id: paperId, action, @@ -784,13 +782,12 @@ export default function ResearchPageNew() { - +
diff --git a/web/src/components/research/SavedPapersList.tsx b/web/src/components/research/SavedPapersList.tsx index 64dd6726..d544bf09 100644 --- a/web/src/components/research/SavedPapersList.tsx +++ b/web/src/components/research/SavedPapersList.tsx @@ -1,6 +1,7 @@ "use client" import { useCallback, useEffect, useMemo, useState } from "react" +import { useSession } from "next-auth/react" import Link from "next/link" import { Check, ChevronDown, Copy, Download, FileText, Filter, Loader2 } from "lucide-react" @@ -308,6 +309,7 @@ function SavedPaperListItem({ } export default function SavedPapersList() { + const { data: session } = useSession() const [items, setItems] = useState([]) const [sortBy, setSortBy] = useState("saved_at") const [page, setPage] = useState(1) @@ -363,7 +365,7 @@ export default function SavedPapersList() { // Fetch tracks on mount useEffect(() => { - fetch("/api/research/tracks?user_id=default", { cache: "no-store" }) + fetch(`/api/research/tracks`, { cache: "no-store" }) .then((res) => res.json()) .then((data) => setTracks(data.tracks || [])) .catch(() => setTracks([])) @@ -376,7 +378,6 @@ export default function SavedPapersList() { const qs = new URLSearchParams({ sort_by: sortBy, limit: "500", - user_id: "default", }) if (selectedTrackId) { qs.set("track_id", String(selectedTrackId)) @@ -452,7 +453,6 @@ export default function SavedPapersList() { setError(null) try { const payload = { - user_id: "default", track_id: activeTrackId, paper_id: externalId || String(paperId), action: requestAction, @@ -500,7 +500,6 @@ export default function SavedPapersList() { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ - user_id: "default", status: nextStatus, }), }) @@ -537,7 +536,7 @@ export default function SavedPapersList() { ) const handleExport = useCallback(async (format: "bibtex" | "ris" | "markdown" | "csl_json") => { - const qs = new URLSearchParams({ format, user_id: "default" }) + const qs = new URLSearchParams({ format }) // Add selected paper IDs selectedIds.forEach((id) => qs.append("paper_id", String(id))) try { @@ -566,7 +565,6 @@ export default function SavedPapersList() { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ - user_id: "default", topic: rwTopic.trim(), }), }) @@ -597,7 +595,6 @@ export default function SavedPapersList() { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ - user_id: "default", content: importBibtex, track_id: selectedTrackId ?? undefined, track_name: selectedTrackId ? undefined : (importTrackName.trim() || undefined), @@ -635,7 +632,6 @@ export default function SavedPapersList() { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ - user_id: "default", track_id: selectedTrackId ?? undefined, track_name: selectedTrackId ? undefined : (zoteroTrackName.trim() || undefined), library_type: zoteroLibraryType, @@ -683,7 +679,7 @@ export default function SavedPapersList() { setCollectionsLoading(true) setCollectionsMessage(null) try { - const qs = new URLSearchParams({ user_id: "default", limit: "200" }) + const qs = new URLSearchParams({ limit: "200" }) if (selectedTrackId) qs.set("track_id", String(selectedTrackId)) const res = await fetch(`/api/research/collections?${qs.toString()}`, { cache: "no-store" }) if (!res.ok) throw new Error(`${res.status}`) @@ -708,7 +704,7 @@ export default function SavedPapersList() { setCollectionsLoading(true) setCollectionsMessage(null) try { - const qs = new URLSearchParams({ user_id: "default", limit: "500" }) + const qs = new URLSearchParams({ limit: "500" }) const res = await fetch(`/api/research/collections/${collectionId}/items?${qs.toString()}`, { cache: "no-store", }) @@ -733,7 +729,6 @@ export default function SavedPapersList() { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ - user_id: "default", name: newCollectionName.trim(), description: newCollectionDesc.trim(), track_id: selectedTrackId ?? undefined, @@ -762,7 +757,6 @@ export default function SavedPapersList() { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ - user_id: "default", paper_id: String(paperId), note: "", tags: [], @@ -793,7 +787,6 @@ export default function SavedPapersList() { method: "PATCH", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ - user_id: "default", note, tags, }), @@ -845,9 +838,8 @@ export default function SavedPapersList() { setCollectionsLoading(true) setCollectionsMessage(null) try { - const qs = new URLSearchParams({ user_id: "default" }) const res = await fetch( - `/api/research/collections/${selectedCollectionId}/items/${paperId}?${qs.toString()}`, + `/api/research/collections/${selectedCollectionId}/items/${paperId}`, { method: "DELETE" }, ) if (!res.ok) throw new Error(`${res.status}`) diff --git a/web/src/components/research/SavedTab.tsx b/web/src/components/research/SavedTab.tsx index 0624fec9..536e64c8 100644 --- a/web/src/components/research/SavedTab.tsx +++ b/web/src/components/research/SavedTab.tsx @@ -47,7 +47,6 @@ type SavedResponse = { } interface SavedTabProps { - userId: string trackId: number | null trackName?: string | null } @@ -67,7 +66,7 @@ function toPaper(item: SavedItem): Paper { } } -export function SavedTab({ userId, trackId, trackName = null }: SavedTabProps) { +export function SavedTab({ trackId, trackName = null }: SavedTabProps) { const [items, setItems] = useState([]) const [loading, setLoading] = useState(false) const [error, setError] = useState(null) @@ -88,10 +87,9 @@ export function SavedTab({ userId, trackId, trackName = null }: SavedTabProps) { buildObsidianExportCommand({ vaultPath: obsidianVaultPath, rootDir: obsidianRootDir, - userId, trackId, }), - [obsidianRootDir, obsidianVaultPath, trackId, userId] + [obsidianRootDir, obsidianVaultPath, trackId] ) const obsidianScope = useMemo( () => describeObsidianScope(trackId, trackName), @@ -99,10 +97,10 @@ export function SavedTab({ userId, trackId, trackName = null }: SavedTabProps) { ) const handleExport = async (format: "bibtex" | "ris" | "markdown" | "csl_json") => { - const qs = new URLSearchParams({ format, user_id: userId }) + const qs = new URLSearchParams({ format }) if (trackId != null) qs.set("track_id", String(trackId)) try { - const res = await fetch(`/api/papers/export?${qs.toString()}`) + const res = await fetch(`/api/papers/export`) if (!res.ok) throw new Error(`${res.status}`) const blob = await res.blob() const extMap: Record = { bibtex: "bib", ris: "ris", markdown: "md", csl_json: "csl.json" } @@ -127,7 +125,6 @@ export function SavedTab({ userId, trackId, trackName = null }: SavedTabProps) { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ - user_id: userId, track_id: trackId, topic: rwTopic.trim(), }), @@ -167,14 +164,13 @@ export function SavedTab({ userId, trackId, trackName = null }: SavedTabProps) { setError(null) try { const qs = new URLSearchParams({ - user_id: userId, sort_by: "saved_at", limit: "100", }) if (trackId != null) { qs.set("track_id", String(trackId)) } - const res = await fetch(`/api/research/papers/saved?${qs.toString()}`) + const res = await fetch(`/api/research/papers/saved`) if (!res.ok) { throw new Error(`${res.status} ${res.statusText}`) } @@ -191,7 +187,7 @@ export function SavedTab({ userId, trackId, trackName = null }: SavedTabProps) { useEffect(() => { load().catch(() => {}) // eslint-disable-next-line react-hooks/exhaustive-deps - }, [userId, trackId]) + }, [trackId]) return (
diff --git a/web/src/components/scholars/ScholarsWatchlist.tsx b/web/src/components/scholars/ScholarsWatchlist.tsx index 72e23dbd..ead6d3f3 100644 --- a/web/src/components/scholars/ScholarsWatchlist.tsx +++ b/web/src/components/scholars/ScholarsWatchlist.tsx @@ -163,7 +163,7 @@ export function ScholarsWatchlist({ scholars }: ScholarsWatchlistProps) { useEffect(() => { const loadTracks = async () => { try { - const res = await fetch("/api/research/tracks?user_id=default", { cache: "no-store" }) + const res = await fetch("/api/research/tracks", { cache: "no-store" }) if (!res.ok) return const payload = (await res.json()) as { tracks?: ResearchTrackSummary[] } setTracks(payload.tracks || []) @@ -395,7 +395,6 @@ export function ScholarsWatchlist({ scholars }: ScholarsWatchlistProps) { method: "POST", headers: { "Content-Type": "application/json" }, body: JSON.stringify({ - user_id: "default", name: trackName, description: `Auto-generated from scholar ${scholar.name}`, keywords, diff --git a/web/src/hooks/useContextPackGeneration.ts b/web/src/hooks/useContextPackGeneration.ts index 04aa6d16..613106ce 100644 --- a/web/src/hooks/useContextPackGeneration.ts +++ b/web/src/hooks/useContextPackGeneration.ts @@ -48,7 +48,6 @@ export function useContextPackGeneration() { try { const payload: Record = { paper_id: params.paperId, - user_id: params.userId ?? "default", depth: params.depth ?? "standard", } if (params.title !== undefined) payload.title = params.title diff --git a/web/src/lib/api.ts b/web/src/lib/api.ts index 3320ea24..41d505cd 100644 --- a/web/src/lib/api.ts +++ b/web/src/lib/api.ts @@ -12,8 +12,7 @@ import { LLMUsageSummary, DeadlineRadarItem, } from "./types" - -const API_BASE_URL = (process.env.PAPERBOT_API_BASE_URL || "http://127.0.0.1:8000") + "/api" +import { API_BASE_URL } from "./config" function slugToName(slug: string): string { return slug @@ -61,7 +60,7 @@ export async function fetchStats(): Promise { } } -export async function fetchActivities(): Promise { +export async function fetchActivities(accessToken?: string): Promise { const activities: Activity[] = [] try { // Fetch recent harvest runs @@ -85,8 +84,10 @@ export async function fetchActivities(): Promise { } catch { /* keep going */ } try { - // Fetch recent saved papers - const savedRes = await fetch(`${API_BASE_URL}/research/papers/saved?user_id=default&limit=3`, { cache: "no-store" }) + // Fetch recent saved papers for the authenticated user + const headers: Record = {} + if (accessToken) headers["Authorization"] = `Bearer ${accessToken}` + const savedRes = await fetch(`${API_BASE_URL}/research/papers/saved?limit=3`, { cache: "no-store", headers }) if (savedRes.ok) { const savedData = await savedRes.json() as { papers?: Array<{ paper_id: string; title: string; authors?: string[]; venue?: string; year?: number; saved_at?: string }> } for (const paper of (savedData.papers || []).slice(0, 3)) { @@ -116,9 +117,11 @@ export async function fetchActivities(): Promise { return activities } -export async function fetchTrendingTopics(): Promise { +export async function fetchTrendingTopics(accessToken?: string): Promise { try { - const res = await fetch(`${API_BASE_URL}/research/tracks?user_id=default`, { cache: "no-store" }) + const headers: Record = {} + if (accessToken) headers["Authorization"] = `Bearer ${accessToken}` + const res = await fetch(`${API_BASE_URL}/research/tracks`, { cache: "no-store", headers }) if (!res.ok) return [] const data = await res.json() as { tracks?: Array<{ keywords?: string[] }> } const keywordCounts = new Map() @@ -157,9 +160,11 @@ export async function fetchPipelineTasks(): Promise { } } -export async function fetchReadingQueue(): Promise { +export async function fetchReadingQueue(accessToken?: string): Promise { try { - const res = await fetch(`${API_BASE_URL}/research/papers/saved?user_id=default&limit=5`, { cache: "no-store" }) + const headers: Record = {} + if (accessToken) headers["Authorization"] = `Bearer ${accessToken}` + const res = await fetch(`${API_BASE_URL}/research/papers/saved?limit=5`, { cache: "no-store", headers }) if (!res.ok) return [] const data = await res.json() as { papers?: Array<{ paper_id: string; title: string }> } return (data.papers || []).slice(0, 5).map((p, i) => ({ @@ -197,16 +202,14 @@ export async function fetchLLMUsage(days: number = 7): Promise } } -export async function fetchDeadlineRadar(userId: string = "default"): Promise { +export async function fetchDeadlineRadar(accessToken?: string): Promise { try { - const qs = new URLSearchParams({ - user_id: userId, - days: "180", - ccf_levels: "A,B,C", - limit: "10", - }) + const qs = new URLSearchParams({ days: "180", ccf_levels: "A,B,C", limit: "10" }) + const headers: Record = {} + if (accessToken) headers["Authorization"] = `Bearer ${accessToken}` const res = await fetch(`${API_BASE_URL}/research/deadlines/radar?${qs.toString()}`, { cache: "no-store", + headers, }) if (!res.ok) return [] const payload = await res.json() as { items?: DeadlineRadarItem[] } @@ -274,7 +277,7 @@ export async function fetchScholars(): Promise { } } -export async function fetchPaperDetails(id: string): Promise { +export async function fetchPaperDetails(id: string, accessToken?: string): Promise { type PaperDetailPayload = { detail?: { paper?: { @@ -300,9 +303,11 @@ export async function fetchPaperDetails(id: string): Promise { } try { + const headers: Record = {} + if (accessToken) headers["Authorization"] = `Bearer ${accessToken}` const res = await fetch( - `${API_BASE_URL}/research/papers/${encodeURIComponent(id)}?user_id=default`, - { cache: "no-store" }, + `${API_BASE_URL}/research/papers/${encodeURIComponent(id)}`, + { cache: "no-store", headers }, ) if (!res.ok) throw new Error("paper detail unavailable") const data = await res.json() as PaperDetailPayload @@ -597,9 +602,11 @@ export async function fetchWikiConcepts(query?: string): Promise } } -export async function fetchPapers(): Promise { +export async function fetchPapers(accessToken?: string): Promise { try { - const res = await fetch(`${API_BASE_URL}/papers/library`) + const headers: Record = {} + if (accessToken) headers["Authorization"] = `Bearer ${accessToken}` + const res = await fetch(`${API_BASE_URL}/papers/library`, { headers }) if (!res.ok) { return [] } diff --git a/web/src/lib/config.ts b/web/src/lib/config.ts new file mode 100644 index 00000000..76b2ede6 --- /dev/null +++ b/web/src/lib/config.ts @@ -0,0 +1,2 @@ +export const API_BASE_URL = + (process.env.PAPERBOT_API_BASE_URL || "http://127.0.0.1:8000") + "/api" diff --git a/web/src/lib/dashboard-api.ts b/web/src/lib/dashboard-api.ts index 7054cb84..b8d5bb22 100644 --- a/web/src/lib/dashboard-api.ts +++ b/web/src/lib/dashboard-api.ts @@ -25,9 +25,11 @@ function formatDateLabel(value?: string | null): string { return parsed.toLocaleDateString("en-US", { month: "short", day: "numeric" }) } -async function fetchJsonOrNull(path: string): Promise { +async function fetchJsonOrNull(path: string, accessToken?: string): Promise { try { - const res = await fetch(`${API_BASE_URL}${path}`, { cache: "no-store" }) + const headers: Record = {} + if (accessToken) headers["Authorization"] = `Bearer ${accessToken}` + const res = await fetch(`${API_BASE_URL}${path}`, { cache: "no-store", headers }) if (!res.ok) return null return (await res.json()) as T } catch { @@ -35,27 +37,24 @@ async function fetchJsonOrNull(path: string): Promise { } } -export async function fetchDashboardTracks(userId: string = "default"): Promise { +export async function fetchDashboardTracks(accessToken?: string): Promise { const payload = await fetchJsonOrNull<{ tracks?: ResearchTrackSummary[] }>( - `/research/tracks?user_id=${encodeURIComponent(userId)}` + `/research/tracks`, + accessToken, ) return payload?.tracks || [] } export async function fetchDashboardTrackFeed( trackId: number, - userId: string = "default", + accessToken?: string, limit: number = 6, ): Promise<{ items: TrackFeedItem[]; total: number }> { - const qs = new URLSearchParams({ - user_id: userId, - limit: String(limit), - offset: "0", - }) + const qs = new URLSearchParams({ limit: String(limit), offset: "0" }) const payload = await fetchJsonOrNull<{ items?: TrackFeedItem[]; total?: number }>( - `/research/tracks/${encodeURIComponent(String(trackId))}/feed?${qs.toString()}` + `/research/tracks/${encodeURIComponent(String(trackId))}/feed?${qs.toString()}`, + accessToken, ) - return { items: payload?.items || [], total: Number(payload?.total || 0), @@ -64,29 +63,26 @@ export async function fetchDashboardTrackFeed( export async function fetchDashboardAnchors( trackId: number, - userId: string = "default", + accessToken?: string, limit: number = 4, ): Promise { const qs = new URLSearchParams({ - user_id: userId, limit: String(limit), window_days: "730", personalized: "true", }) const payload = await fetchJsonOrNull<{ items?: AnchorPreviewItem[] }>( - `/research/tracks/${encodeURIComponent(String(trackId))}/anchors/discover?${qs.toString()}` + `/research/tracks/${encodeURIComponent(String(trackId))}/anchors/discover?${qs.toString()}`, + accessToken, ) return payload?.items || [] } export async function fetchDashboardReadingQueue( - userId: string = "default", + accessToken?: string, limit: number = 8, ): Promise { - const qs = new URLSearchParams({ - user_id: userId, - limit: String(limit), - }) + const qs = new URLSearchParams({ limit: String(limit) }) const payload = await fetchJsonOrNull<{ items?: Array<{ saved_at?: string @@ -97,7 +93,7 @@ export async function fetchDashboardReadingQueue( authors?: string[] } }> - }>(`/research/papers/saved?${qs.toString()}`) + }>(`/research/papers/saved?${qs.toString()}`, accessToken) return (payload?.items || []).map((row, index) => { const paper = row.paper || {} @@ -113,12 +109,11 @@ export async function fetchDashboardReadingQueue( } export async function fetchIntelligenceFeed( - userId: string = "default", + accessToken?: string, limit: number = 6, filters?: IntelligenceFeedFilters, ): Promise { const qs = new URLSearchParams({ - user_id: userId, limit: String(limit), }) if (filters?.source) qs.set("source", filters.source) @@ -128,7 +123,7 @@ export async function fetchIntelligenceFeed( if (filters?.sortOrder) qs.set("sort_order", filters.sortOrder) if (filters?.trackId) qs.set("track_id", String(filters.trackId)) - const payload = await fetchJsonOrNull(`/intelligence/feed?${qs.toString()}`) + const payload = await fetchJsonOrNull(`/intelligence/feed?${qs.toString()}`, accessToken) return payload || { items: [], refreshed_at: null, @@ -139,7 +134,7 @@ export async function fetchIntelligenceFeed( } } -export async function fetchDashboardActivities(userId: string = "default"): Promise { +export async function fetchDashboardActivities(accessToken?: string): Promise { const [runsPayload, savedPayload] = await Promise.all([ fetchJsonOrNull<{ runs?: Array<{ @@ -161,7 +156,7 @@ export async function fetchDashboardActivities(userId: string = "default"): Prom year?: number | null } }> - }>(`/research/papers/saved?user_id=${encodeURIComponent(userId)}&limit=4`), + }>(`/research/papers/saved?limit=4`, accessToken), ]) const activities: Activity[] = [] diff --git a/web/src/middleware.ts b/web/src/middleware.ts new file mode 100644 index 00000000..24e4b81e --- /dev/null +++ b/web/src/middleware.ts @@ -0,0 +1,36 @@ +import { auth } from "@/auth" +import { NextResponse } from "next/server" + +// Only allow unauthenticated access to explicit auth pages. +// All other paths (including "/") require a valid session. +const PUBLIC_PATHS = ["/login", "/register", "/forgot-password", "/reset-password"] + +export default auth((req) => { + const { pathname } = req.nextUrl + + // Skip all Next internals and API routes + if ( + pathname.startsWith("/_next") || + pathname.startsWith("/api/") || + pathname === "/favicon.ico" + ) { + return NextResponse.next() + } + + // Public pages + if (PUBLIC_PATHS.some((p) => pathname === p || pathname.startsWith(p + "/"))) { + return NextResponse.next() + } + + if (!req.auth) { + const url = new URL("/login", req.url) + url.searchParams.set("callbackUrl", req.nextUrl.pathname) + return NextResponse.redirect(url) + } + + return NextResponse.next() +}) + +export const config = { + matcher: ["/(.*)"], +} diff --git a/web/src/types/next-auth.d.ts b/web/src/types/next-auth.d.ts new file mode 100644 index 00000000..57991f51 --- /dev/null +++ b/web/src/types/next-auth.d.ts @@ -0,0 +1,9 @@ +import "next-auth" + +declare module "next-auth" { + interface Session { + accessToken?: string + userId?: number + provider?: string + } +} From 6c804026ee0bf560898f4871a074d3cb34b877b2 Mon Sep 17 00:00:00 2001 From: WenjingWang Date: Fri, 13 Mar 2026 00:22:43 +0100 Subject: [PATCH 7/9] feat(web): add track context proxy route --- .../api/research/tracks/[trackId]/context/route.ts | 13 +++++++++++++ 1 file changed, 13 insertions(+) create mode 100644 web/src/app/api/research/tracks/[trackId]/context/route.ts diff --git a/web/src/app/api/research/tracks/[trackId]/context/route.ts b/web/src/app/api/research/tracks/[trackId]/context/route.ts new file mode 100644 index 00000000..1ecb8cd9 --- /dev/null +++ b/web/src/app/api/research/tracks/[trackId]/context/route.ts @@ -0,0 +1,13 @@ +export const runtime = "nodejs" + +import { apiBaseUrl, proxyJson } from "../../../_base" + +export async function GET(req: Request, ctx: { params: Promise<{ trackId: string }> }) { + const { trackId } = await ctx.params + const url = new URL(req.url) + return proxyJson( + req, + `${apiBaseUrl()}/api/research/tracks/${encodeURIComponent(trackId)}/context?${url.searchParams.toString()}`, + "GET", + ) +} From 8a84b02d4cebf7a6cd432a5da528444d8d396a27 Mon Sep 17 00:00:00 2001 From: jerry <1772030600@qq.com> Date: Fri, 13 Mar 2026 14:21:12 +0800 Subject: [PATCH 8/9] fix(research): unblock global feedback and scholars --- .../versions/0027_global_paper_feedback.py | 51 ++++++++ src/paperbot/agents/__init__.py | 52 ++++---- src/paperbot/api/routes/research.py | 69 +++++++---- src/paperbot/infrastructure/stores/models.py | 30 +++-- .../infrastructure/stores/research_store.py | 93 +++++++++++++-- tests/unit/test_research_feedback_state.py | 111 ++++++++++++++++++ tests/unit/test_research_scholar_routes.py | 17 +++ web/src/components/research/PaperCard.tsx | 16 ++- .../components/research/ResearchPageNew.tsx | 107 ++++++++++------- 9 files changed, 438 insertions(+), 108 deletions(-) create mode 100644 alembic/versions/0027_global_paper_feedback.py diff --git a/alembic/versions/0027_global_paper_feedback.py b/alembic/versions/0027_global_paper_feedback.py new file mode 100644 index 00000000..9c4a1ae4 --- /dev/null +++ b/alembic/versions/0027_global_paper_feedback.py @@ -0,0 +1,51 @@ +"""Allow global paper feedback without an active track. + +Revision ID: 0027_global_paper_feedback +Revises: 0026_embedding_endpoint_settings +Create Date: 2026-03-13 +""" + +from __future__ import annotations + +from alembic import op +import sqlalchemy as sa + + +revision = "0027_global_paper_feedback" +down_revision = "0026_embedding_endpoint_settings" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + conn = op.get_bind() + inspector = sa.inspect(conn) + if "paper_feedback" not in inspector.get_table_names(): + return + + columns = {column["name"]: column for column in inspector.get_columns("paper_feedback")} + track_column = columns.get("track_id") + if not track_column or bool(track_column.get("nullable")): + return + + with op.batch_alter_table("paper_feedback") as batch_op: + batch_op.alter_column("track_id", existing_type=sa.Integer(), nullable=True) + + +def downgrade() -> None: + conn = op.get_bind() + inspector = sa.inspect(conn) + if "paper_feedback" not in inspector.get_table_names(): + return + + columns = {column["name"]: column for column in inspector.get_columns("paper_feedback")} + track_column = columns.get("track_id") + if not track_column or not bool(track_column.get("nullable")): + return + + null_count = conn.execute(sa.text("SELECT COUNT(*) FROM paper_feedback WHERE track_id IS NULL")).scalar() + if null_count: + raise RuntimeError("Cannot downgrade: global paper_feedback rows with NULL track_id exist") + + with op.batch_alter_table("paper_feedback") as batch_op: + batch_op.alter_column("track_id", existing_type=sa.Integer(), nullable=False) diff --git a/src/paperbot/agents/__init__.py b/src/paperbot/agents/__init__.py index 3fa217b3..9aecc241 100644 --- a/src/paperbot/agents/__init__.py +++ b/src/paperbot/agents/__init__.py @@ -1,26 +1,10 @@ -# src/paperbot/agents/__init__.py -""" -PaperBot Agent 模块。 - -提供各类 AI Agent 实现: -- BaseAgent: Agent 基类 -- ResearchAgent: 论文研究 Agent -- CodeAnalysisAgent: 代码分析 Agent -- QualityAgent: 质量评估 Agent -- DocumentationAgent: 文档生成 Agent -- ConferenceResearchAgent: 会议论文抓取 Agent -- ReviewerAgent: 论文评审 Agent -- VerificationAgent: 声明验证 Agent +"""PaperBot agent exports. + +Keep imports lazy so routes that only need a lightweight scholar agent do not +pull in optional report/PDF dependencies during package initialization. """ -from .base import BaseAgent -from .research.agent import ResearchAgent -from .code_analysis.agent import CodeAnalysisAgent -from .quality.agent import QualityAgent -from .documentation.agent import DocumentationAgent -from .conference.agent import ConferenceResearchAgent -from .review.agent import ReviewerAgent -from .verification.agent import VerificationAgent +from importlib import import_module __all__ = [ "BaseAgent", @@ -32,3 +16,29 @@ "ReviewerAgent", "VerificationAgent", ] + +_LAZY_EXPORTS = { + "BaseAgent": (".base", "BaseAgent"), + "ResearchAgent": (".research.agent", "ResearchAgent"), + "CodeAnalysisAgent": (".code_analysis.agent", "CodeAnalysisAgent"), + "QualityAgent": (".quality.agent", "QualityAgent"), + "DocumentationAgent": (".documentation.agent", "DocumentationAgent"), + "ConferenceResearchAgent": (".conference.agent", "ConferenceResearchAgent"), + "ReviewerAgent": (".review.agent", "ReviewerAgent"), + "VerificationAgent": (".verification.agent", "VerificationAgent"), +} + + +def __getattr__(name: str): + try: + module_name, attr_name = _LAZY_EXPORTS[name] + except KeyError as exc: + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") from exc + + value = getattr(import_module(module_name, __name__), attr_name) + globals()[name] = value + return value + + +def __dir__() -> list[str]: + return sorted(set(globals()) | set(__all__)) diff --git a/src/paperbot/api/routes/research.py b/src/paperbot/api/routes/research.py index 463ad0d9..2abc07ad 100644 --- a/src/paperbot/api/routes/research.py +++ b/src/paperbot/api/routes/research.py @@ -1164,12 +1164,15 @@ def add_paper_feedback( track_id = req.track_id active_track: Optional[Dict[str, Any]] = None if track_id is None: - Logger.info("No track specified, getting active track", file=LogFiles.HARVEST) + Logger.info("No track specified, checking active track", file=LogFiles.HARVEST) active_track = research_store.get_active_track(user_id=user_id) - if not active_track: - Logger.error("No active track found", file=LogFiles.HARVEST) - raise HTTPException(status_code=400, detail="track_id missing and no active track") - track_id = int(active_track["id"]) + if active_track: + track_id = int(active_track["id"]) + else: + Logger.info( + "No active track found; recording global paper feedback", + file=LogFiles.HARVEST, + ) meta: Dict[str, Any] = dict(req.metadata or {}) if req.context_run_id is not None: @@ -1243,7 +1246,7 @@ def add_paper_feedback( Logger.info("Paper feedback recorded successfully", file=LogFiles.HARVEST) normalized_action = research_store._normalize_feedback_action(req.action) current_action = research_store._effective_feedback_action(normalized_action) - if normalized_action in {"save", "unsave"}: + if normalized_action in {"save", "unsave"} and track_id is not None: export_track = active_track or research_store.get_track_by_id(track_id=int(track_id)) _schedule_obsidian_export_for_track( background_tasks, @@ -1429,7 +1432,9 @@ def list_saved_papers( @router.post("/research/discovery/seed", response_model=DiscoverySeedResponse) -async def discover_from_seed(req: DiscoverySeedRequest, user_id: str = Depends(get_required_user_id)): +async def discover_from_seed( + req: DiscoverySeedRequest, user_id: str = Depends(get_required_user_id) +): if req.year_from and req.year_to and req.year_from > req.year_to: raise HTTPException(status_code=400, detail="year_from must be <= year_to") @@ -1608,9 +1613,7 @@ async def discover_from_seed(req: DiscoverySeedRequest, user_id: str = Depends(g year_to=req.year_to, ) feedback_profile = ( - _build_feedback_profile(user_id=user_id, track_id=req.track_id) - if req.personalized - else {} + _build_feedback_profile(user_id=user_id, track_id=req.track_id) if req.personalized else {} ) scored = _rank_discovery_candidates( filtered, @@ -1680,7 +1683,9 @@ async def discover_from_seed(req: DiscoverySeedRequest, user_id: str = Depends(g @router.post("/research/collections", response_model=PaperCollectionResponse) -def create_collection(req: PaperCollectionCreateRequest, user_id: str = Depends(get_required_user_id)): +def create_collection( + req: PaperCollectionCreateRequest, user_id: str = Depends(get_required_user_id) +): try: collection = _get_research_store().create_collection( user_id=user_id, @@ -1710,7 +1715,11 @@ def list_collections( @router.patch("/research/collections/{collection_id}", response_model=PaperCollectionResponse) -def update_collection(collection_id: int, req: PaperCollectionUpdateRequest, user_id: str = Depends(get_required_user_id)): +def update_collection( + collection_id: int, + req: PaperCollectionUpdateRequest, + user_id: str = Depends(get_required_user_id), +): try: collection = _get_research_store().update_collection( user_id=user_id, @@ -1747,7 +1756,11 @@ def list_collection_items( "/research/collections/{collection_id}/items", response_model=PaperCollectionItemsResponse, ) -def upsert_collection_item(collection_id: int, req: PaperCollectionItemUpsertRequest, user_id: str = Depends(get_required_user_id)): +def upsert_collection_item( + collection_id: int, + req: PaperCollectionItemUpsertRequest, + user_id: str = Depends(get_required_user_id), +): item = _get_research_store().upsert_collection_item( user_id=user_id, collection_id=collection_id, @@ -1757,17 +1770,22 @@ def upsert_collection_item(collection_id: int, req: PaperCollectionItemUpsertReq ) if item is None: raise HTTPException(status_code=404, detail="Collection or paper not found") - items = _get_research_store().list_collection_items(user_id=user_id, collection_id=collection_id) - return PaperCollectionItemsResponse( - user_id=user_id, collection_id=collection_id, items=items + items = _get_research_store().list_collection_items( + user_id=user_id, collection_id=collection_id ) + return PaperCollectionItemsResponse(user_id=user_id, collection_id=collection_id, items=items) @router.patch( "/research/collections/{collection_id}/items/{paper_id}", response_model=PaperCollectionItemsResponse, ) -def patch_collection_item(collection_id: int, paper_id: str, req: PaperCollectionItemPatchRequest, user_id: str = Depends(get_required_user_id)): +def patch_collection_item( + collection_id: int, + paper_id: str, + req: PaperCollectionItemPatchRequest, + user_id: str = Depends(get_required_user_id), +): item = _get_research_store().upsert_collection_item( user_id=user_id, collection_id=collection_id, @@ -1777,14 +1795,16 @@ def patch_collection_item(collection_id: int, paper_id: str, req: PaperCollectio ) if item is None: raise HTTPException(status_code=404, detail="Collection or paper not found") - items = _get_research_store().list_collection_items(user_id=user_id, collection_id=collection_id) - return PaperCollectionItemsResponse( - user_id=user_id, collection_id=collection_id, items=items + items = _get_research_store().list_collection_items( + user_id=user_id, collection_id=collection_id ) + return PaperCollectionItemsResponse(user_id=user_id, collection_id=collection_id, items=items) @router.delete("/research/collections/{collection_id}/items/{paper_id}") -def delete_collection_item(collection_id: int, paper_id: str, user_id: str = Depends(get_required_user_id)): +def delete_collection_item( + collection_id: int, paper_id: str, user_id: str = Depends(get_required_user_id) +): ok = _get_research_store().remove_collection_item( user_id=user_id, collection_id=collection_id, @@ -2223,7 +2243,12 @@ def discover_track_anchors( "/research/tracks/{track_id}/anchors/{author_id}/action", response_model=AnchorActionResponse, ) -def set_anchor_action(track_id: int, author_id: int, req: AnchorActionRequest, user_id: str = Depends(get_required_user_id)): +def set_anchor_action( + track_id: int, + author_id: int, + req: AnchorActionRequest, + user_id: str = Depends(get_required_user_id), +): _ensure_anchor_feature_enabled() track = _get_research_store().get_track(user_id=user_id, track_id=track_id) diff --git a/src/paperbot/infrastructure/stores/models.py b/src/paperbot/infrastructure/stores/models.py index 89780844..fc79e697 100644 --- a/src/paperbot/infrastructure/stores/models.py +++ b/src/paperbot/infrastructure/stores/models.py @@ -490,14 +490,19 @@ class ResearchMilestoneModel(Base): class PaperFeedbackModel(Base): - """User feedback on recommended/seen papers (track-scoped).""" + """User feedback on recommended/seen papers (track-scoped or global).""" __tablename__ = "paper_feedback" id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) user_id: Mapped[str] = mapped_column(String(64), index=True) - track_id: Mapped[int] = mapped_column(Integer, ForeignKey("research_tracks.id"), index=True) + track_id: Mapped[Optional[int]] = mapped_column( + Integer, + ForeignKey("research_tracks.id"), + nullable=True, + index=True, + ) paper_id: Mapped[str] = mapped_column(String(64), index=True) paper_ref_id: Mapped[Optional[int]] = mapped_column( @@ -1055,9 +1060,7 @@ class DocumentIndexJobModel(Base): error: Mapped[Optional[str]] = mapped_column(Text, nullable=True) enqueued_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), index=True) started_at: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True) - finished_at: Mapped[Optional[datetime]] = mapped_column( - DateTime(timezone=True), nullable=True - ) + finished_at: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True) updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), index=True) paper = relationship("PaperModel") @@ -1297,7 +1300,9 @@ class ReproContextPackModel(Base): paper_id: Mapped[str] = mapped_column(String(256), nullable=False, index=True) paper_title: Mapped[Optional[str]] = mapped_column(Text, nullable=True) version: Mapped[str] = mapped_column(String(16), nullable=False, default="v1") - depth: Mapped[str] = mapped_column(String(16), nullable=False, default="standard") # fast/standard/deep + depth: Mapped[str] = mapped_column( + String(16), nullable=False, default="standard" + ) # fast/standard/deep status: Mapped[str] = mapped_column( String(32), nullable=False, default="pending", index=True ) # pending/running/completed/failed @@ -1340,7 +1345,9 @@ class ReproContextStageResultModel(Base): String(64), ForeignKey("repro_context_pack.id", ondelete="CASCADE"), index=True ) stage_name: Mapped[str] = mapped_column(String(64), index=True) - status: Mapped[str] = mapped_column(String(16), default="completed", index=True) # completed/failed/skipped + status: Mapped[str] = mapped_column( + String(16), default="completed", index=True + ) # completed/failed/skipped result_json: Mapped[Optional[str]] = mapped_column(Text, nullable=True) confidence: Mapped[float] = mapped_column(Float, default=0.0) duration_ms: Mapped[int] = mapped_column(Integer, default=0) @@ -1359,7 +1366,9 @@ class ReproContextEvidenceModel(Base): context_pack_id: Mapped[str] = mapped_column( String(64), ForeignKey("repro_context_pack.id", ondelete="CASCADE"), index=True ) - evidence_type: Mapped[str] = mapped_column(String(32), index=True) # paper_span/table/figure/code_snippet/metadata + evidence_type: Mapped[str] = mapped_column( + String(32), index=True + ) # paper_span/table/figure/code_snippet/metadata ref: Mapped[str] = mapped_column(Text, default="") supports_json: Mapped[str] = mapped_column(Text, default="[]") # JSON array of field names confidence: Mapped[float] = mapped_column(Float, default=0.0) @@ -1427,6 +1436,7 @@ class ReproCodeExperienceModel(Base): code_snippet: Mapped[Optional[str]] = mapped_column(Text, nullable=True) created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), index=True) + class IntelligenceEventModel(Base): """Cached community radar signal from external sources.""" @@ -1497,7 +1507,9 @@ class UserModel(Base): avatar_url: Mapped[Optional[str]] = mapped_column(String(512), nullable=True) is_active: Mapped[bool] = mapped_column(Boolean, nullable=False, default=True) created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) - last_login_at: Mapped[Optional[datetime]] = mapped_column(DateTime(timezone=True), nullable=True) + last_login_at: Mapped[Optional[datetime]] = mapped_column( + DateTime(timezone=True), nullable=True + ) # ============================================================================ diff --git a/src/paperbot/infrastructure/stores/research_store.py b/src/paperbot/infrastructure/stores/research_store.py index f141e417..2bcc966e 100644 --- a/src/paperbot/infrastructure/stores/research_store.py +++ b/src/paperbot/infrastructure/stores/research_store.py @@ -6,7 +6,7 @@ from datetime import datetime, timedelta, timezone from typing import Any, Dict, Iterable, List, Optional -from sqlalchemy import desc, func, or_, select +from sqlalchemy import desc, func, inspect, or_, select from sqlalchemy.exc import IntegrityError from paperbot.application.ports.feedback_port import FeedbackPort @@ -81,6 +81,7 @@ def _parse_datetime(value: Any) -> Optional[datetime]: except Exception: return None + _FEEDBACK_ACTION_ALIASES: Dict[str, str] = { "not_relevant": "dislike", "not-relevant": "dislike", @@ -106,6 +107,10 @@ def _parse_datetime(value: Any) -> Optional[datetime]: "skip": "skip", "cite": "cite", } + +_LEGACY_GLOBAL_FEEDBACK_TRACK_NAME = "__paperbot_legacy_global_feedback__" + + class SqlAlchemyResearchStore(FeedbackPort): """ Track/progress store for personalized paper recommendation. @@ -163,6 +168,54 @@ def _feedback_state_key( cls._feedback_group(normalized_action), ) + @staticmethod + def _paper_feedback_track_is_nullable(session) -> bool: + try: + columns = inspect(session.bind).get_columns("paper_feedback") + except Exception: + return True + + for column in columns: + if str(column.get("name")) == "track_id": + return bool(column.get("nullable")) + return True + + def _ensure_legacy_global_feedback_track( + self, + *, + session, + user_id: str, + now: datetime, + ) -> int: + row = session.execute( + select(ResearchTrackModel).where( + ResearchTrackModel.user_id == user_id, + ResearchTrackModel.name == _LEGACY_GLOBAL_FEEDBACK_TRACK_NAME, + ) + ).scalar_one_or_none() + if row is None: + row = ResearchTrackModel( + user_id=user_id, + name=_LEGACY_GLOBAL_FEEDBACK_TRACK_NAME, + description="System track for legacy global feedback compatibility.", + keywords_json="[]", + venues_json="[]", + methods_json="[]", + is_active=0, + archived_at=now, + created_at=now, + updated_at=now, + ) + session.add(row) + session.flush() + else: + row.is_active = 0 + row.archived_at = row.archived_at or now + row.updated_at = now + session.add(row) + + return int(row.id) + @classmethod def _collapse_effective_feedback_rows( cls, rows: Iterable[PaperFeedbackModel] @@ -491,7 +544,7 @@ def add_paper_feedback( self, *, user_id: str, - track_id: int, + track_id: Optional[int], paper_id: str, action: str, weight: float = 0.0, @@ -502,14 +555,29 @@ def add_paper_feedback( metadata = dict(metadata or {}) normalized_action = self._normalize_feedback_action(action) with self._provider.session() as session: - track = session.execute( - select(ResearchTrackModel).where( - ResearchTrackModel.user_id == user_id, ResearchTrackModel.id == track_id + track = None + resolved_track_id: Optional[int] = None + if track_id is not None: + track = session.execute( + select(ResearchTrackModel).where( + ResearchTrackModel.user_id == user_id, ResearchTrackModel.id == track_id + ) + ).scalar_one_or_none() + if track is None: + Logger.error("Track not found", file=LogFiles.HARVEST) + return None + resolved_track_id = int(track_id) + elif not self._paper_feedback_track_is_nullable(session): + Logger.info( + "paper_feedback.track_id is still NOT NULL; using legacy global feedback track", + file=LogFiles.HARVEST, ) - ).scalar_one_or_none() - if track is None: - Logger.error("Track not found", file=LogFiles.HARVEST) - return None + resolved_track_id = self._ensure_legacy_global_feedback_track( + session=session, + user_id=user_id, + now=now, + ) + metadata.setdefault("global_feedback_mode", "legacy_track_fallback") resolved_paper_ref_id = self._resolve_paper_ref_id( session=session, @@ -519,7 +587,7 @@ def add_paper_feedback( Logger.info("Creating new feedback record", file=LogFiles.HARVEST) row = PaperFeedbackModel( user_id=user_id, - track_id=track_id, + track_id=resolved_track_id, paper_id=(paper_id or "").strip(), paper_ref_id=resolved_paper_ref_id, canonical_paper_id=resolved_paper_ref_id, # dual-write @@ -562,8 +630,9 @@ def add_paper_feedback( now=now, ) - track.updated_at = now - session.add(track) + if track is not None: + track.updated_at = now + session.add(track) session.commit() session.refresh(row) Logger.info("Feedback record created successfully", file=LogFiles.HARVEST) diff --git a/tests/unit/test_research_feedback_state.py b/tests/unit/test_research_feedback_state.py index d711e2a6..c4f6a0f0 100644 --- a/tests/unit/test_research_feedback_state.py +++ b/tests/unit/test_research_feedback_state.py @@ -3,8 +3,10 @@ from pathlib import Path from fastapi.testclient import TestClient +from sqlalchemy import text from paperbot.api import main as api_main +from paperbot.api.auth.dependencies import get_required_user_id from paperbot.api.routes import research as research_route from paperbot.infrastructure.stores.paper_store import SqlAlchemyPaperStore from paperbot.infrastructure.stores.research_store import SqlAlchemyResearchStore @@ -216,3 +218,112 @@ def get_track_by_id(self, *, track_id: int): assert captured == [ {"user_id": "persisted-owner", "track_id": 9, "for_tracks": False}, ] + + +def test_global_feedback_without_active_track_is_accepted(monkeypatch): + class _FakeResearchStore: + def get_active_track(self, *, user_id: str): + assert user_id == "u-feedback" + return None + + def add_paper_feedback(self, **kwargs): + assert kwargs["user_id"] == "u-feedback" + assert kwargs["track_id"] is None + assert kwargs["paper_id"] == "paper-1" + return {"id": 1, "track_id": None, "paper_id": "paper-1", "action": "save"} + + def _normalize_feedback_action(self, action: str) -> str: + return action + + def _effective_feedback_action(self, action: str): + return action + + monkeypatch.setattr(research_route, "_research_store", _FakeResearchStore()) + + captured: list[dict[str, object]] = [] + monkeypatch.setattr( + research_route, + "_schedule_obsidian_export_for_track", + lambda *args, **kwargs: captured.append({"called": True}), + ) + + api_main.app.dependency_overrides[get_required_user_id] = lambda: "u-feedback" + try: + with TestClient(api_main.app) as client: + response = client.post( + "/api/research/papers/feedback", + json={ + "paper_id": "paper-1", + "action": "save", + "paper_title": "Global Save", + }, + ) + finally: + api_main.app.dependency_overrides.pop(get_required_user_id, None) + + assert response.status_code == 200 + assert response.json()["current_action"] == "save" + assert captured == [] + + +def test_store_allows_global_feedback_without_track(tmp_path: Path): + store, paper, _ = _prepare_feedback_state_db(tmp_path) + paper_id = str(paper["id"]) + + feedback = store.add_paper_feedback( + user_id="u-feedback", + track_id=None, + paper_id=paper_id, + action="save", + metadata={"title": paper["title"]}, + ) + + assert feedback is not None + assert feedback["track_id"] is None + + saved = store.list_saved_papers(user_id="u-feedback", limit=10) + assert len(saved) == 1 + assert str(saved[0]["id"]) == paper_id + + +def test_store_falls_back_when_legacy_schema_still_requires_track(tmp_path: Path): + store, paper, _ = _prepare_feedback_state_db(tmp_path) + paper_id = str(paper["id"]) + + with store._provider.engine.begin() as conn: + conn.execute(text("DROP TABLE paper_feedback")) + conn.execute(text(""" + CREATE TABLE paper_feedback ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + user_id VARCHAR(64), + track_id INTEGER NOT NULL, + paper_id VARCHAR(64), + paper_ref_id INTEGER, + action VARCHAR(32), + canonical_paper_id INTEGER, + weight FLOAT, + ts DATETIME, + metadata_json TEXT, + FOREIGN KEY(track_id) REFERENCES research_tracks (id), + FOREIGN KEY(paper_ref_id) REFERENCES papers (id), + FOREIGN KEY(canonical_paper_id) REFERENCES papers (id) + ) + """)) + + feedback = store.add_paper_feedback( + user_id="u-feedback", + track_id=None, + paper_id=paper_id, + action="save", + metadata={"title": paper["title"]}, + ) + + assert feedback is not None + assert feedback["track_id"] is not None + + visible_tracks = store.list_tracks(user_id="u-feedback") + assert all(track["name"] != "__paperbot_legacy_global_feedback__" for track in visible_tracks) + + saved = store.list_saved_papers(user_id="u-feedback", limit=10) + assert len(saved) == 1 + assert str(saved[0]["id"]) == paper_id diff --git a/tests/unit/test_research_scholar_routes.py b/tests/unit/test_research_scholar_routes.py index bc608a73..4e2272b5 100644 --- a/tests/unit/test_research_scholar_routes.py +++ b/tests/unit/test_research_scholar_routes.py @@ -1,3 +1,6 @@ +import importlib +import sys + from fastapi.testclient import TestClient from paperbot.api import main as api_main @@ -300,3 +303,17 @@ def test_scholar_search_route(monkeypatch): assert payload["query"] == "alice" assert payload["total"] == 2 assert payload["items"][0]["author_id"] == "1001" + + +def test_scholar_profile_agent_import_is_lazy(): + for module_name in [ + "paperbot.agents", + "paperbot.agents.review.agent", + "paperbot.agents.scholar_tracking.scholar_profile_agent", + ]: + sys.modules.pop(module_name, None) + + module = importlib.import_module("paperbot.agents.scholar_tracking.scholar_profile_agent") + + assert hasattr(module, "ScholarProfileAgent") + assert "paperbot.agents.review.agent" not in sys.modules diff --git a/web/src/components/research/PaperCard.tsx b/web/src/components/research/PaperCard.tsx index 34f9997d..f965e368 100644 --- a/web/src/components/research/PaperCard.tsx +++ b/web/src/components/research/PaperCard.tsx @@ -119,8 +119,12 @@ export function PaperCard({ const requestAction = toggleSaveFeedbackAction(isSaved) setActionLoading("save") try { - await onFeedbackAction(requestAction) - setIsSaved((prev) => !prev) + const nextAction = await onFeedbackAction(requestAction) + if (nextAction !== undefined) { + setIsSaved(nextAction === "save") + } + } catch { + // Parent surface handles the user-facing error state. } finally { setActionLoading(null) } @@ -131,8 +135,12 @@ export function PaperCard({ const requestAction = togglePaperPreferenceAction(preferenceAction, targetAction) setActionLoading(targetAction) try { - await onFeedbackAction(requestAction) - setPreferenceAction(requestAction === targetAction ? targetAction : null) + const nextAction = await onFeedbackAction(requestAction) + if (nextAction !== undefined) { + setPreferenceAction(normalizePaperPreferenceAction(nextAction)) + } + } catch { + // Parent surface handles the user-facing error state. } finally { setActionLoading(null) } diff --git a/web/src/components/research/ResearchPageNew.tsx b/web/src/components/research/ResearchPageNew.tsx index dba3c1bd..a56f41e0 100644 --- a/web/src/components/research/ResearchPageNew.tsx +++ b/web/src/components/research/ResearchPageNew.tsx @@ -1,7 +1,6 @@ "use client" import { useEffect, useMemo, useState } from "react" -import { useSession } from "next-auth/react" import Link from "next/link" import { useSearchParams } from "next/navigation" @@ -58,6 +57,9 @@ type ContextPack = { paper_recommendations?: Paper[] paper_recommendation_reasons?: Record } + +const RESULT_LIMIT_OPTIONS = [10, 25, 50] as const + function getGreeting(): string { const hour = new Date().getHours() if (hour < 12) return "Good morning" @@ -66,7 +68,6 @@ function getGreeting(): string { } export default function ResearchPageNew() { - const { data: session } = useSession() const searchParams = useSearchParams() // User state @@ -86,6 +87,7 @@ export default function ResearchPageNew() { const [isSearching, setIsSearching] = useState(false) const [contextPack, setContextPack] = useState(null) const [searchSources, setSearchSources] = useState(ALL_SOURCES) + const [paperLimit, setPaperLimit] = useState<(typeof RESULT_LIMIT_OPTIONS)[number]>(25) const [yearFrom, setYearFrom] = useState("") const [yearTo, setYearTo] = useState("") @@ -127,7 +129,6 @@ export default function ResearchPageNew() { // Load tracks on mount useEffect(() => { refreshTracks().catch((e) => setError(getErrorMessage(e))) - // eslint-disable-next-line react-hooks/exhaustive-deps }, []) useEffect(() => { @@ -241,7 +242,7 @@ export default function ResearchPageNew() { const body = { query, track_id: activeTrackId ?? undefined, - paper_limit: 10, + paper_limit: paperLimit, memory_limit: 8, sources: searchSources, offline: false, @@ -292,7 +293,7 @@ export default function ResearchPageNew() { return () => clearTimeout(timer) // eslint-disable-next-line react-hooks/exhaustive-deps - }, [searchSources, anchorPersonalized, yearFrom, yearTo]) + }, [searchSources, anchorPersonalized, paperLimit, yearFrom, yearTo]) // If the page is opened with a query parameter, run it once automatically. useEffect(() => { @@ -430,44 +431,54 @@ export default function ResearchPageNew() { action: PaperFeedbackRequestAction, rank?: number, paper?: Paper - ): Promise { - // Don't set global loading - PaperCard handles its own loading state + ): Promise { setError(null) - const body: Record = { - track_id: activeTrackId, - paper_id: paperId, - action, - weight: 0.0, - context_run_id: contextPack?.context_run_id ?? null, - context_rank: typeof rank === "number" ? rank : undefined, - metadata: { - retrieval_sources: Array.isArray(paper?.retrieval_sources) - ? paper?.retrieval_sources - : [], - retrieval_score: - typeof paper?.retrieval_score === "number" ? paper.retrieval_score : undefined, - anchor_mode: anchorPersonalized ? "personalized" : "global", - }, - } + try { + const body: Record = { + track_id: activeTrackId, + paper_id: paperId, + action, + weight: 0.0, + context_run_id: contextPack?.context_run_id ?? null, + context_rank: typeof rank === "number" ? rank : undefined, + metadata: { + retrieval_sources: Array.isArray(paper?.retrieval_sources) + ? paper?.retrieval_sources + : [], + retrieval_score: + typeof paper?.retrieval_score === "number" ? paper.retrieval_score : undefined, + anchor_mode: anchorPersonalized ? "personalized" : "global", + }, + } - // Include paper metadata for save action - if (action === "save" && paper) { - body.paper_title = paper.title - body.paper_abstract = paper.abstract || "" - body.paper_authors = paper.authors || [] - body.paper_year = paper.year - body.paper_venue = paper.venue - body.paper_citation_count = paper.citation_count - body.paper_url = paper.url - body.paper_source = paper.source || "semantic_scholar" - } + if (action === "save" && paper) { + body.paper_title = paper.title + body.paper_abstract = paper.abstract || "" + body.paper_authors = paper.authors || [] + body.paper_year = paper.year + body.paper_venue = paper.venue + body.paper_citation_count = paper.citation_count + body.paper_url = paper.url + body.paper_source = paper.source || "semantic_scholar" + } - const payload = await fetchJson<{ current_action?: string | null }>(`/api/research/papers/feedback`, { - method: "POST", - body: JSON.stringify(body), - headers: { "Content-Type": "application/json" }, - }) - return normalizePaperFeedbackAction(payload.current_action) ?? currentFeedbackFromRequestAction(action) + const payload = await fetchJson<{ current_action?: string | null }>( + `/api/research/papers/feedback`, + { + method: "POST", + body: JSON.stringify(body), + headers: { "Content-Type": "application/json" }, + } + ) + + return ( + normalizePaperFeedbackAction(payload.current_action) ?? + currentFeedbackFromRequestAction(action) + ) + } catch (e) { + setError(getErrorMessage(e)) + return undefined + } } const trackToClearName = tracks.find((t) => t.id === trackToClear)?.name || "this track" @@ -692,6 +703,22 @@ export default function ResearchPageNew() { Sources: {searchSources.length} Results: {papers.length} +
+ Cap + {RESULT_LIMIT_OPTIONS.map((limit) => ( + + ))} +
-
- - setUserId(e.target.value)} className="w-[200px]" /> - -
+
{error ? (