diff --git a/alembic/versions/0007_paper_harvest_tables.py b/alembic/versions/0007_paper_harvest_tables.py index 6677a8de..09746ce6 100644 --- a/alembic/versions/0007_paper_harvest_tables.py +++ b/alembic/versions/0007_paper_harvest_tables.py @@ -42,12 +42,21 @@ def _get_indexes(table: str) -> set[str]: return idx +def _get_columns(table: str) -> set[str]: + cols = set() + for c in _insp().get_columns(table): + cols.add(str(c.get("name") or "")) + return cols + + def _create_index(name: str, table: str, cols: list[str]) -> None: if _is_offline(): op.create_index(name, table, cols) return if name in _get_indexes(table): return + if not set(cols).issubset(_get_columns(table)): + return op.create_index(name, table, cols) diff --git a/alembic/versions/0010_contract_feedback_fk.py b/alembic/versions/0010_contract_feedback_fk.py new file mode 100644 index 00000000..0973a467 --- /dev/null +++ b/alembic/versions/0010_contract_feedback_fk.py @@ -0,0 +1,97 @@ +"""contract feedback fk read path + +Revision ID: 0010_contract_feedback_fk +Revises: 0009_paper_identifiers +Create Date: 2026-02-12 08:05:00 + +Contract phase for canonical feedback join: +- Backfill canonical_paper_id from paper_ref_id when available. +- Add index optimized for library reads by canonical FK. +- Drop legacy paper_id index used by external-id joins. +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import context, op + + +revision = "0010_contract_feedback_fk" +down_revision = "0009_paper_identifiers" +branch_labels = None +depends_on = None + + +def _is_offline() -> bool: + try: + return bool(context.is_offline_mode()) + except Exception: + return False + + +def _inspector(): + bind = op.get_bind() + return sa.inspect(bind) + + +def _has_table(name: str) -> bool: + if _is_offline(): + return False + return bool(_inspector().has_table(name)) + + +def _get_columns(table: str) -> set[str]: + if _is_offline() or not _has_table(table): + return set() + return {c["name"] for c in _inspector().get_columns(table)} + + +def _has_index(table: str, index_name: str) -> bool: + if _is_offline() or not _has_table(table): + return False + names = {idx.get("name") for idx in _inspector().get_indexes(table)} + return index_name in names + + +def _create_index(name: str, table: str, cols: list[str]) -> None: + if _is_offline() or _has_index(table, name): + return + op.create_index(name, table, cols) + + +def _drop_index(name: str, table: str) -> None: + if _is_offline() or not _has_index(table, name): + return + op.drop_index(name, table_name=table) + + +def upgrade() -> None: + if not _has_table("paper_feedback"): + return + + cols = _get_columns("paper_feedback") + if {"canonical_paper_id", "paper_ref_id"}.issubset(cols): + op.execute( + sa.text( + """ + UPDATE paper_feedback + SET canonical_paper_id = paper_ref_id + WHERE canonical_paper_id IS NULL + AND paper_ref_id IS NOT NULL + """ + ) + ) + + _create_index( + "ix_paper_feedback_user_action_canonical", + "paper_feedback", + ["user_id", "action", "canonical_paper_id"], + ) + + # Legacy external-id join path index. + _drop_index("ix_paper_feedback_paper_id", "paper_feedback") + + +def downgrade() -> None: + _create_index("ix_paper_feedback_paper_id", "paper_feedback", ["paper_id"]) + _drop_index("ix_paper_feedback_user_action_canonical", "paper_feedback") diff --git a/alembic/versions/0011_model_endpoints.py b/alembic/versions/0011_model_endpoints.py new file mode 100644 index 00000000..034ef42c --- /dev/null +++ b/alembic/versions/0011_model_endpoints.py @@ -0,0 +1,72 @@ +"""add model endpoint gateway table + +Revision ID: 0011_model_endpoints +Revises: 0010_contract_feedback_fk +Create Date: 2026-02-12 11:10:00 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import context, op + + +revision = "0011_model_endpoints" +down_revision = "0010_contract_feedback_fk" +branch_labels = None +depends_on = None + + +def _is_offline() -> bool: + try: + return bool(context.is_offline_mode()) + except Exception: + return False + + +def _inspector(): + bind = op.get_bind() + return sa.inspect(bind) + + +def _has_table(name: str) -> bool: + if _is_offline(): + return False + return bool(_inspector().has_table(name)) + + +def upgrade() -> None: + if _has_table("model_endpoints"): + return + + op.create_table( + "model_endpoints", + sa.Column("id", sa.Integer(), primary_key=True, autoincrement=True), + sa.Column("name", sa.String(length=64), nullable=False), + sa.Column( + "vendor", sa.String(length=32), nullable=False, server_default="openai_compatible" + ), + sa.Column("base_url", sa.String(length=512), nullable=True), + sa.Column( + "api_key_env", sa.String(length=64), nullable=False, server_default="OPENAI_API_KEY" + ), + sa.Column("models_json", sa.Text(), nullable=False, server_default="[]"), + sa.Column("task_types_json", sa.Text(), nullable=False, server_default="[]"), + sa.Column("enabled", sa.Boolean(), nullable=False, server_default=sa.text("1")), + sa.Column("is_default", sa.Boolean(), nullable=False, server_default=sa.text("0")), + sa.Column("created_at", sa.DateTime(timezone=True), nullable=False), + sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False), + sa.UniqueConstraint("name", name="uq_model_endpoints_name"), + ) + op.create_index("ix_model_endpoints_name", "model_endpoints", ["name"]) + op.create_index("ix_model_endpoints_enabled", "model_endpoints", ["enabled"]) + op.create_index("ix_model_endpoints_is_default", "model_endpoints", ["is_default"]) + + +def downgrade() -> None: + if not _has_table("model_endpoints"): + return + op.drop_index("ix_model_endpoints_is_default", table_name="model_endpoints") + op.drop_index("ix_model_endpoints_enabled", table_name="model_endpoints") + op.drop_index("ix_model_endpoints_name", table_name="model_endpoints") + op.drop_table("model_endpoints") diff --git a/alembic/versions/0012_reconcile_papers_schema.py b/alembic/versions/0012_reconcile_papers_schema.py new file mode 100644 index 00000000..d01e47e0 --- /dev/null +++ b/alembic/versions/0012_reconcile_papers_schema.py @@ -0,0 +1,134 @@ +"""reconcile papers schema with current ORM + +Revision ID: 0012_reconcile_papers_schema +Revises: 0011_model_endpoints +Create Date: 2026-02-12 12:20:00 + +Adds missing columns/indexes on legacy `papers` table created before 0007. +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import context, op + + +revision = "0012_reconcile_papers_schema" +down_revision = "0011_model_endpoints" +branch_labels = None +depends_on = None + + +def _is_offline() -> bool: + try: + return bool(context.is_offline_mode()) + except Exception: + return False + + +def _inspector(): + return sa.inspect(op.get_bind()) + + +def _has_table(name: str) -> bool: + if _is_offline(): + return False + return bool(_inspector().has_table(name)) + + +def _columns(table: str) -> set[str]: + if _is_offline() or not _has_table(table): + return set() + return {str(c.get("name") or "") for c in _inspector().get_columns(table)} + + +def _has_index(table: str, index_name: str) -> bool: + if _is_offline() or not _has_table(table): + return False + names = {str(i.get("name") or "") for i in _inspector().get_indexes(table)} + return index_name in names + + +def _add_column_if_missing(table: str, column: sa.Column) -> None: + if _is_offline() or column.name in _columns(table): + return + op.add_column(table, column) + + +def _create_index_if_possible(name: str, table: str, cols: list[str]) -> None: + if _is_offline() or _has_index(table, name): + return + if not set(cols).issubset(_columns(table)): + return + op.create_index(name, table, cols) + + +def upgrade() -> None: + if not _has_table("papers"): + return + + _add_column_if_missing( + "papers", sa.Column("semantic_scholar_id", sa.String(length=64), nullable=True) + ) + _add_column_if_missing("papers", sa.Column("openalex_id", sa.String(length=64), nullable=True)) + _add_column_if_missing( + "papers", + sa.Column("title_hash", sa.String(length=64), nullable=False, server_default=""), + ) + _add_column_if_missing("papers", sa.Column("year", sa.Integer(), nullable=True)) + _add_column_if_missing( + "papers", sa.Column("publication_date", sa.String(length=32), nullable=True) + ) + _add_column_if_missing( + "papers", sa.Column("citation_count", sa.Integer(), nullable=False, server_default="0") + ) + _add_column_if_missing( + "papers", + sa.Column("fields_of_study_json", sa.Text(), nullable=False, server_default="[]"), + ) + _add_column_if_missing( + "papers", + sa.Column("primary_source", sa.String(length=32), nullable=False, server_default=""), + ) + _add_column_if_missing( + "papers", + sa.Column("sources_json", sa.Text(), nullable=False, server_default="[]"), + ) + _add_column_if_missing( + "papers", sa.Column("deleted_at", sa.DateTime(timezone=True), nullable=True) + ) + + # Fill empty title_hash so ORM non-null assumptions hold. + if "title_hash" in _columns("papers"): + op.execute( + sa.text( + """ + UPDATE papers + SET title_hash = lower(hex(randomblob(16))) + WHERE title_hash IS NULL OR title_hash = '' + """ + ) + ) + + _create_index_if_possible("ix_papers_semantic_scholar_id", "papers", ["semantic_scholar_id"]) + _create_index_if_possible("ix_papers_openalex_id", "papers", ["openalex_id"]) + _create_index_if_possible("ix_papers_title_hash", "papers", ["title_hash"]) + _create_index_if_possible("ix_papers_year", "papers", ["year"]) + _create_index_if_possible("ix_papers_venue", "papers", ["venue"]) + _create_index_if_possible("ix_papers_citation_count", "papers", ["citation_count"]) + _create_index_if_possible("ix_papers_primary_source", "papers", ["primary_source"]) + + +def downgrade() -> None: + # Keep columns for backward compatibility; only drop indexes created here. + for idx in [ + "ix_papers_semantic_scholar_id", + "ix_papers_openalex_id", + "ix_papers_title_hash", + "ix_papers_year", + "ix_papers_venue", + "ix_papers_citation_count", + "ix_papers_primary_source", + ]: + if _has_index("papers", idx): + op.drop_index(idx, table_name="papers") diff --git a/alembic/versions/0013_model_endpoint_api_key.py b/alembic/versions/0013_model_endpoint_api_key.py new file mode 100644 index 00000000..7efff54b --- /dev/null +++ b/alembic/versions/0013_model_endpoint_api_key.py @@ -0,0 +1,56 @@ +"""add api key storage for model endpoints + +Revision ID: 0013_model_endpoint_api_key +Revises: 0012_reconcile_papers_schema +Create Date: 2026-02-12 20:30:00 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import context, op + + +revision = "0013_model_endpoint_api_key" +down_revision = "0012_reconcile_papers_schema" +branch_labels = None +depends_on = None + + +def _is_offline() -> bool: + try: + return bool(context.is_offline_mode()) + except Exception: + return False + + +def _inspector(): + return sa.inspect(op.get_bind()) + + +def _has_table(name: str) -> bool: + if _is_offline(): + return False + return bool(_inspector().has_table(name)) + + +def _columns(table: str) -> set[str]: + if _is_offline() or not _has_table(table): + return set() + return {str(c.get("name") or "") for c in _inspector().get_columns(table)} + + +def upgrade() -> None: + if not _has_table("model_endpoints"): + return + if "api_key_value" not in _columns("model_endpoints"): + op.add_column( + "model_endpoints", sa.Column("api_key_value", sa.String(length=512), nullable=True) + ) + + +def downgrade() -> None: + if not _has_table("model_endpoints"): + return + if "api_key_value" in _columns("model_endpoints"): + op.drop_column("model_endpoints", "api_key_value") diff --git a/alembic/versions/0014_llm_usage.py b/alembic/versions/0014_llm_usage.py new file mode 100644 index 00000000..2929e3e2 --- /dev/null +++ b/alembic/versions/0014_llm_usage.py @@ -0,0 +1,86 @@ +"""add llm usage tracking table + +Revision ID: 0014_llm_usage +Revises: 0013_model_endpoint_api_key +Create Date: 2026-02-12 21:25:00 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import context, op + + +revision = "0014_llm_usage" +down_revision = "0013_model_endpoint_api_key" +branch_labels = None +depends_on = None + + +def _is_offline() -> bool: + try: + return bool(context.is_offline_mode()) + except Exception: + return False + + +def _inspector(): + return sa.inspect(op.get_bind()) + + +def _has_table(name: str) -> bool: + if _is_offline(): + return False + return bool(_inspector().has_table(name)) + + +def _has_index(table: str, index_name: str) -> bool: + if _is_offline() or not _has_table(table): + return False + names = {str(i.get("name") or "") for i in _inspector().get_indexes(table)} + return index_name in names + + +def upgrade() -> None: + if _has_table("llm_usage"): + return + + op.create_table( + "llm_usage", + sa.Column("id", sa.Integer(), primary_key=True, autoincrement=True), + sa.Column("ts", sa.DateTime(timezone=True), nullable=False), + sa.Column("task_type", sa.String(length=32), nullable=False, server_default="default"), + sa.Column("provider_name", sa.String(length=64), nullable=False, server_default="unknown"), + sa.Column("model_name", sa.String(length=128), nullable=False, server_default=""), + sa.Column("prompt_tokens", sa.Integer(), nullable=False, server_default="0"), + sa.Column("completion_tokens", sa.Integer(), nullable=False, server_default="0"), + sa.Column("total_tokens", sa.Integer(), nullable=False, server_default="0"), + sa.Column("estimated_cost_usd", sa.Float(), nullable=False, server_default="0"), + sa.Column("metadata_json", sa.Text(), nullable=False, server_default="{}"), + ) + + if not _has_index("llm_usage", "ix_llm_usage_ts"): + op.create_index("ix_llm_usage_ts", "llm_usage", ["ts"]) + if not _has_index("llm_usage", "ix_llm_usage_task_type"): + op.create_index("ix_llm_usage_task_type", "llm_usage", ["task_type"]) + if not _has_index("llm_usage", "ix_llm_usage_provider_name"): + op.create_index("ix_llm_usage_provider_name", "llm_usage", ["provider_name"]) + if not _has_index("llm_usage", "ix_llm_usage_model_name"): + op.create_index("ix_llm_usage_model_name", "llm_usage", ["model_name"]) + if not _has_index("llm_usage", "ix_llm_usage_total_tokens"): + op.create_index("ix_llm_usage_total_tokens", "llm_usage", ["total_tokens"]) + + +def downgrade() -> None: + if not _has_table("llm_usage"): + return + for idx in [ + "ix_llm_usage_total_tokens", + "ix_llm_usage_model_name", + "ix_llm_usage_provider_name", + "ix_llm_usage_task_type", + "ix_llm_usage_ts", + ]: + if _has_index("llm_usage", idx): + op.drop_index(idx, table_name="llm_usage") + op.drop_table("llm_usage") diff --git a/alembic/versions/0015_pipeline_sessions.py b/alembic/versions/0015_pipeline_sessions.py new file mode 100644 index 00000000..650d942f --- /dev/null +++ b/alembic/versions/0015_pipeline_sessions.py @@ -0,0 +1,87 @@ +"""add pipeline session checkpoints table + +Revision ID: 0015_pipeline_sessions +Revises: 0014_llm_usage +Create Date: 2026-02-12 22:10:00 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import context, op + + +revision = "0015_pipeline_sessions" +down_revision = "0014_llm_usage" +branch_labels = None +depends_on = None + + +def _is_offline() -> bool: + try: + return bool(context.is_offline_mode()) + except Exception: + return False + + +def _inspector(): + return sa.inspect(op.get_bind()) + + +def _has_table(name: str) -> bool: + if _is_offline(): + return False + return bool(_inspector().has_table(name)) + + +def _has_index(table: str, index_name: str) -> bool: + if _is_offline() or not _has_table(table): + return False + names = {str(i.get("name") or "") for i in _inspector().get_indexes(table)} + return index_name in names + + +def upgrade() -> None: + if _has_table("pipeline_sessions"): + return + + op.create_table( + "pipeline_sessions", + sa.Column("session_id", sa.String(length=64), primary_key=True), + sa.Column("workflow", sa.String(length=64), nullable=False, server_default=""), + sa.Column("status", sa.String(length=32), nullable=False, server_default="running"), + sa.Column("checkpoint", sa.String(length=64), nullable=False, server_default="init"), + sa.Column("payload_json", sa.Text(), nullable=False, server_default="{}"), + sa.Column("state_json", sa.Text(), nullable=False, server_default="{}"), + sa.Column("result_json", sa.Text(), nullable=False, server_default="{}"), + sa.Column("created_at", sa.DateTime(timezone=True), nullable=False), + sa.Column("updated_at", sa.DateTime(timezone=True), nullable=False), + ) + + if not _has_index("pipeline_sessions", "ix_pipeline_sessions_workflow"): + op.create_index("ix_pipeline_sessions_workflow", "pipeline_sessions", ["workflow"]) + if not _has_index("pipeline_sessions", "ix_pipeline_sessions_status"): + op.create_index("ix_pipeline_sessions_status", "pipeline_sessions", ["status"]) + if not _has_index("pipeline_sessions", "ix_pipeline_sessions_checkpoint"): + op.create_index("ix_pipeline_sessions_checkpoint", "pipeline_sessions", ["checkpoint"]) + if not _has_index("pipeline_sessions", "ix_pipeline_sessions_created_at"): + op.create_index("ix_pipeline_sessions_created_at", "pipeline_sessions", ["created_at"]) + if not _has_index("pipeline_sessions", "ix_pipeline_sessions_updated_at"): + op.create_index("ix_pipeline_sessions_updated_at", "pipeline_sessions", ["updated_at"]) + + +def downgrade() -> None: + if not _has_table("pipeline_sessions"): + return + + for idx in [ + "ix_pipeline_sessions_updated_at", + "ix_pipeline_sessions_created_at", + "ix_pipeline_sessions_checkpoint", + "ix_pipeline_sessions_status", + "ix_pipeline_sessions_workflow", + ]: + if _has_index("pipeline_sessions", idx): + op.drop_index(idx, table_name="pipeline_sessions") + + op.drop_table("pipeline_sessions") diff --git a/scripts/backfill_identifiers.py b/scripts/backfill_identifiers.py index 4620e156..1c8f0d69 100644 --- a/scripts/backfill_identifiers.py +++ b/scripts/backfill_identifiers.py @@ -9,6 +9,7 @@ from __future__ import annotations import argparse +import json import sys from datetime import datetime, timezone from pathlib import Path @@ -19,6 +20,7 @@ sys.path.insert(0, str(SRC)) from sqlalchemy import select # noqa: E402 +from paperbot.application.services.identity_resolver import IdentityResolver # noqa: E402 from paperbot.infrastructure.stores.models import ( # noqa: E402 Base, PaperFeedbackModel, @@ -76,26 +78,51 @@ def backfill_identifiers(provider: SessionProvider) -> dict: return {"identifiers_created": created, "identifiers_skipped": skipped} -def backfill_canonical_paper_id(provider: SessionProvider) -> dict: - """Populate paper_feedback.canonical_paper_id from paper_ref_id.""" +def backfill_canonical_paper_id(provider: SessionProvider, db_url: str) -> dict: + """Populate paper_feedback.canonical_paper_id from paper_ref_id / IdentityResolver.""" + resolver = IdentityResolver(db_url=db_url) updated = 0 + resolved_from_ref = 0 + resolved_from_identity = 0 + unresolved = 0 with provider.session() as session: rows = ( session.execute( select(PaperFeedbackModel).where( PaperFeedbackModel.canonical_paper_id.is_(None), - PaperFeedbackModel.paper_ref_id.is_not(None), ) ) .scalars() .all() ) for row in rows: - row.canonical_paper_id = row.paper_ref_id - updated += 1 + resolved_id = int(row.paper_ref_id) if row.paper_ref_id is not None else None + if resolved_id is None: + try: + metadata = json.loads(row.metadata_json or "{}") + if not isinstance(metadata, dict): + metadata = {} + except Exception: + metadata = {} + resolved_id = resolver.resolve(str(row.paper_id or "").strip(), hints=metadata) + + if resolved_id is not None: + row.canonical_paper_id = int(resolved_id) + updated += 1 + if row.paper_ref_id is not None: + resolved_from_ref += 1 + else: + resolved_from_identity += 1 + else: + unresolved += 1 session.commit() - return {"feedback_rows_updated": updated} + return { + "feedback_rows_updated": updated, + "resolved_from_paper_ref_id": resolved_from_ref, + "resolved_from_identity_resolver": resolved_from_identity, + "feedback_rows_unresolved": unresolved, + } def main() -> None: @@ -112,7 +139,7 @@ def main() -> None: print(result1) print("=== Backfilling canonical_paper_id ===") - result2 = backfill_canonical_paper_id(provider) + result2 = backfill_canonical_paper_id(provider, db_url) print(result2) print("Done.") diff --git a/src/paperbot/api/main.py b/src/paperbot/api/main.py index 3b2e9153..66a0fb3a 100644 --- a/src/paperbot/api/main.py +++ b/src/paperbot/api/main.py @@ -22,6 +22,7 @@ paperscool, newsletter, harvest, + model_endpoints, ) from paperbot.infrastructure.event_log.logging_event_log import LoggingEventLog from paperbot.infrastructure.event_log.composite_event_log import CompositeEventLog @@ -67,6 +68,7 @@ async def health_check(): app.include_router(paperscool.router, prefix="/api", tags=["PapersCool"]) app.include_router(newsletter.router, prefix="/api", tags=["Newsletter"]) app.include_router(harvest.router, prefix="/api", tags=["Harvest"]) +app.include_router(model_endpoints.router, prefix="/api", tags=["Model Endpoints"]) @app.on_event("startup") diff --git a/src/paperbot/api/routes/__init__.py b/src/paperbot/api/routes/__init__.py index 6a1344d9..eef8cb9a 100644 --- a/src/paperbot/api/routes/__init__.py +++ b/src/paperbot/api/routes/__init__.py @@ -14,6 +14,7 @@ research, paperscool, newsletter, + model_endpoints, ) __all__ = [ @@ -30,4 +31,5 @@ "research", "paperscool", "newsletter", + "model_endpoints", ] diff --git a/src/paperbot/api/routes/model_endpoints.py b/src/paperbot/api/routes/model_endpoints.py new file mode 100644 index 00000000..73e15fb6 --- /dev/null +++ b/src/paperbot/api/routes/model_endpoints.py @@ -0,0 +1,202 @@ +from __future__ import annotations + +import os +from typing import Any, Dict, List, Optional + +from fastapi import APIRouter, HTTPException +from pydantic import BaseModel, Field + +from paperbot.infrastructure.llm.router import ModelConfig, ModelRouter, RouterConfig +from paperbot.infrastructure.stores.llm_usage_store import LLMUsageStore +from paperbot.infrastructure.stores.model_endpoint_store import ModelEndpointStore + +router = APIRouter() + +_store = ModelEndpointStore() +_usage_store = LLMUsageStore() + +_ALLOWED_VENDORS = ["openai_compatible", "openai", "anthropic", "ollama"] +_ALLOWED_TASK_TYPES = [ + "default", + "extraction", + "summary", + "analysis", + "reasoning", + "code", + "review", + "chat", +] + + +class ModelEndpointCreateRequest(BaseModel): + name: str = Field(..., min_length=1, max_length=64) + vendor: str = "openai_compatible" + base_url: Optional[str] = None + api_key_env: str = "OPENAI_API_KEY" + api_key: Optional[str] = None + models: List[str] = Field(default_factory=list) + task_types: List[str] = Field(default_factory=list) + enabled: bool = True + is_default: bool = False + + +class ModelEndpointUpdateRequest(BaseModel): + name: Optional[str] = Field(default=None, min_length=1, max_length=64) + vendor: Optional[str] = None + base_url: Optional[str] = None + api_key_env: Optional[str] = None + api_key: Optional[str] = None + models: Optional[List[str]] = None + task_types: Optional[List[str]] = None + enabled: Optional[bool] = None + is_default: Optional[bool] = None + + +class ModelEndpointListResponse(BaseModel): + items: List[Dict[str, Any]] + + +class ModelEndpointResponse(BaseModel): + item: Dict[str, Any] + + +class EndpointTestRequest(BaseModel): + remote: bool = False + + +class EndpointTestResponse(BaseModel): + ok: bool + endpoint_id: int + provider: Dict[str, Any] + api_key_present: bool + message: str + + +class EndpointCapabilitiesResponse(BaseModel): + vendors: List[str] + task_types: List[str] + + +class EndpointActivateResponse(BaseModel): + item: Dict[str, Any] + + +class LLMUsageSummaryResponse(BaseModel): + summary: Dict[str, Any] + + +def _build_model_config(endpoint: Dict[str, Any]) -> ModelConfig: + models = [str(x).strip() for x in (endpoint.get("models") or []) if str(x).strip()] + if not models: + raise ValueError("endpoint has no models") + + vendor = str(endpoint.get("vendor") or "openai_compatible").strip().lower() + provider = ModelRouter._normalize_provider(vendor) + return ModelConfig( + provider=provider, + model=models[0], + api_key_env=str(endpoint.get("api_key_env") or "OPENAI_API_KEY"), + api_key=(str(endpoint.get("api_key") or "").strip() or None), + base_url=(str(endpoint.get("base_url") or "").strip() or None), + ) + + +@router.get("/model-endpoints", response_model=ModelEndpointListResponse) +def list_model_endpoints(enabled_only: bool = False): + rows = _store.list_endpoints(enabled_only=enabled_only) + return ModelEndpointListResponse(items=rows) + + +@router.get("/model-endpoints/capabilities", response_model=EndpointCapabilitiesResponse) +def get_model_endpoint_capabilities(): + return EndpointCapabilitiesResponse(vendors=_ALLOWED_VENDORS, task_types=_ALLOWED_TASK_TYPES) + + +@router.post("/model-endpoints", response_model=ModelEndpointResponse) +def create_model_endpoint(req: ModelEndpointCreateRequest): + try: + row = _store.upsert_endpoint(payload=req.model_dump()) + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + return ModelEndpointResponse(item=row) + + +@router.patch("/model-endpoints/{endpoint_id}", response_model=ModelEndpointResponse) +def update_model_endpoint(endpoint_id: int, req: ModelEndpointUpdateRequest): + existing = _store.get_endpoint(endpoint_id) + if not existing: + raise HTTPException(status_code=404, detail="model endpoint not found") + + payload = {k: v for k, v in req.model_dump().items() if v is not None} + try: + row = _store.upsert_endpoint(payload=payload, endpoint_id=endpoint_id) + except ValueError as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc + return ModelEndpointResponse(item=row) + + +@router.delete("/model-endpoints/{endpoint_id}") +def delete_model_endpoint(endpoint_id: int): + ok = _store.delete_endpoint(endpoint_id) + if not ok: + raise HTTPException(status_code=404, detail="model endpoint not found") + return {"ok": True} + + +@router.post("/model-endpoints/{endpoint_id}/activate", response_model=EndpointActivateResponse) +def activate_model_endpoint(endpoint_id: int): + row = _store.activate_endpoint(endpoint_id) + if not row: + raise HTTPException(status_code=404, detail="model endpoint not found") + return EndpointActivateResponse(item=row) + + +@router.get("/model-endpoints/usage", response_model=LLMUsageSummaryResponse) +def get_llm_usage_summary(days: int = 7): + window = max(1, min(int(days), 90)) + summary = _usage_store.summarize(days=window) + return LLMUsageSummaryResponse(summary=summary) + + +@router.post("/model-endpoints/{endpoint_id}/test", response_model=EndpointTestResponse) +def test_model_endpoint(endpoint_id: int, req: EndpointTestRequest): + endpoint = _store.get_endpoint(endpoint_id, include_secrets=True) + if not endpoint: + raise HTTPException(status_code=404, detail="model endpoint not found") + + api_key_env = str(endpoint.get("api_key_env") or "OPENAI_API_KEY") + api_key_present = bool(str(endpoint.get("api_key") or "").strip()) or bool( + os.getenv(api_key_env) + ) + + try: + cfg = _build_model_config(endpoint) + router = ModelRouter(RouterConfig(models={"__test__": cfg}, fallback_model="__test__")) + provider = router.get_provider("default") + info = provider.info + + message = "Provider initialized successfully." + if req.remote: + if not api_key_present: + raise ValueError(f"missing API key env: {api_key_env}") + pong = provider.invoke_simple( + "You are a connection checker.", + "Reply with a single word: OK", + max_tokens=16, + temperature=0, + ) + message = f"Remote check success: {(pong or '').strip()[:80]}" + + return EndpointTestResponse( + ok=True, + endpoint_id=endpoint_id, + provider={ + "provider_name": info.provider_name, + "model_name": info.model_name, + "api_base": info.api_base, + }, + api_key_present=api_key_present, + message=message, + ) + except Exception as exc: + raise HTTPException(status_code=400, detail=f"test failed: {exc}") from exc diff --git a/src/paperbot/api/routes/paperscool.py b/src/paperbot/api/routes/paperscool.py index d44f20bb..5a09f237 100644 --- a/src/paperbot/api/routes/paperscool.py +++ b/src/paperbot/api/routes/paperscool.py @@ -15,12 +15,18 @@ from paperbot.api.streaming import StreamEvent, wrap_generator from paperbot.application.services.daily_push_service import DailyPushService from paperbot.application.services.llm_service import get_llm_service +from paperbot.application.services.enrichment_pipeline import ( + EnrichmentContext, + EnrichmentPipeline, + FilterStep, + JudgeStep, + LLMEnrichmentStep, +) +from paperbot.application.services.paper_search_service import PaperSearchService from paperbot.application.workflows.analysis.paper_judge import PaperJudge from paperbot.application.workflows.dailypaper import ( DailyPaperReporter, - apply_judge_scores_to_report, build_daily_paper_report, - enrich_daily_paper_report, ingest_daily_report_to_registry, normalize_llm_features, normalize_output_formats, @@ -28,11 +34,20 @@ render_daily_paper_markdown, select_judge_candidates, ) -from paperbot.application.workflows.paperscool_topic_search import PapersCoolTopicSearchWorkflow +from paperbot.application.workflows.unified_topic_search import ( + make_default_search_service, + run_unified_topic_search, +) +from paperbot.infrastructure.stores.paper_store import PaperStore +from paperbot.infrastructure.stores.pipeline_session_store import PipelineSessionStore from paperbot.infrastructure.stores.research_store import SqlAlchemyResearchStore from paperbot.utils.text_processing import extract_github_url router = APIRouter() +_paper_search_service: Optional[PaperSearchService] = None +_pipeline_session_store = PipelineSessionStore() +# Test compatibility hook: unit tests can monkeypatch this to inject a fake workflow. +PapersCoolTopicSearchWorkflow = None _ALLOWED_REPORT_BASE = os.path.abspath("./reports") @@ -62,6 +77,45 @@ def _validate_email_list(emails: List[str]) -> List[str]: return cleaned +def _get_paper_search_service() -> PaperSearchService: + global _paper_search_service + if _paper_search_service is None: + _paper_search_service = make_default_search_service(registry=PaperStore()) + return _paper_search_service + + +async def _run_topic_search( + *, + queries: List[str], + sources: List[str], + branches: List[str], + top_k_per_query: int, + show_per_branch: int, + min_score: float, +) -> Dict[str, Any]: + if callable(PapersCoolTopicSearchWorkflow): + workflow = PapersCoolTopicSearchWorkflow() + return workflow.run( + queries=queries, + sources=sources, + branches=branches, + top_k_per_query=top_k_per_query, + show_per_branch=show_per_branch, + min_score=min_score, + ) + + return await run_unified_topic_search( + queries=queries, + sources=sources, + branches=branches, + top_k_per_query=top_k_per_query, + show_per_branch=show_per_branch, + min_score=min_score, + search_service=_get_paper_search_service(), + persist=False, + ) + + class PapersCoolSearchRequest(BaseModel): queries: List[str] = Field(default_factory=list) sources: List[str] = Field(default_factory=lambda: ["papers_cool"]) @@ -103,6 +157,14 @@ class DailyPaperRequest(BaseModel): notify: bool = False notify_channels: List[str] = Field(default_factory=list) notify_email_to: List[str] = Field(default_factory=list) + session_id: Optional[str] = Field( + default=None, description="Resume token for long-running pipeline" + ) + resume: bool = Field(False, description="Resume from latest persisted checkpoint") + require_approval: bool = Field( + False, + description="Pause before registry ingest and require manual approve/reject", + ) class DailyPaperResponse(BaseModel): @@ -113,6 +175,18 @@ class DailyPaperResponse(BaseModel): notify_result: Optional[Dict[str, Any]] = None +class PipelineSessionResponse(BaseModel): + session: Dict[str, Any] + + +class ApprovalQueueResponse(BaseModel): + items: List[Dict[str, Any]] + + +class ApprovalDecisionRequest(BaseModel): + reason: str = "" + + class PapersCoolAnalyzeRequest(BaseModel): report: Dict[str, Any] run_judge: bool = False @@ -141,14 +215,13 @@ class PapersCoolReposResponse(BaseModel): @router.post("/research/paperscool/search", response_model=PapersCoolSearchResponse) -def topic_search(req: PapersCoolSearchRequest): +async def topic_search(req: PapersCoolSearchRequest): cleaned_queries = [q.strip() for q in req.queries if (q or "").strip()] if not cleaned_queries: raise HTTPException(status_code=400, detail="queries is required") - workflow = PapersCoolTopicSearchWorkflow() try: - result = workflow.run( + result = await _run_topic_search( queries=cleaned_queries, sources=req.sources, branches=req.branches, @@ -165,18 +238,92 @@ async def _dailypaper_stream(req: DailyPaperRequest): """SSE generator for the full DailyPaper pipeline.""" cleaned_queries = [q.strip() for q in req.queries if (q or "").strip()] + session = _pipeline_session_store.start_session( + workflow="paperscool_daily", + payload=req.model_dump(), + session_id=req.session_id, + resume=req.resume, + ) + session_id = str(session.get("session_id") or "") + session_state: Dict[str, Any] = session.get("state") if req.resume else {} + + yield StreamEvent( + type="status", + data={ + "phase": "session", + "session_id": session_id, + "resume": bool(req.resume), + "checkpoint": session.get("checkpoint") or "init", + }, + ) + + if ( + req.resume + and session.get("status") == "completed" + and isinstance(session.get("result"), dict) + ): + cached_result = dict(session.get("result") or {}) + payload = { + "report": cached_result.get("report") or {}, + "markdown": cached_result.get("markdown") or "", + "markdown_path": cached_result.get("markdown_path"), + "json_path": cached_result.get("json_path"), + "notify_result": cached_result.get("notify_result"), + "session_id": session_id, + "resumed": True, + } + yield StreamEvent(type="result", data=payload) + return + + if ( + req.resume + and session.get("status") == "pending_approval" + and isinstance(session.get("result"), dict) + ): + cached_result = dict(session.get("result") or {}) + payload = { + "report": cached_result.get("report") or {}, + "markdown": cached_result.get("markdown") or "", + "markdown_path": cached_result.get("markdown_path"), + "json_path": cached_result.get("json_path"), + "notify_result": cached_result.get("notify_result"), + "session_id": session_id, + "resumed": True, + "approval_status": "pending_approval", + } + yield StreamEvent( + type="approval_required", + data={"phase": "approval", "session_id": session_id, "status": "pending_approval"}, + ) + yield StreamEvent(type="result", data=payload) + return + # Phase 1 — Search - yield StreamEvent(type="progress", data={"phase": "search", "message": "Searching papers..."}) - workflow = PapersCoolTopicSearchWorkflow() effective_top_k = max(int(req.top_k_per_query), int(req.top_n), 1) - search_result = workflow.run( - queries=cleaned_queries, - sources=req.sources, - branches=req.branches, - top_k_per_query=effective_top_k, - show_per_branch=req.show_per_branch, - min_score=req.min_score, - ) + if req.resume and isinstance(session_state.get("search_result"), dict): + search_result = dict(session_state.get("search_result") or {}) + yield StreamEvent( + type="progress", + data={"phase": "search", "message": "Resumed search result from checkpoint"}, + ) + else: + yield StreamEvent( + type="progress", data={"phase": "search", "message": "Searching papers..."} + ) + search_result = await _run_topic_search( + queries=cleaned_queries, + sources=req.sources, + branches=req.branches, + top_k_per_query=effective_top_k, + show_per_branch=req.show_per_branch, + min_score=req.min_score, + ) + _pipeline_session_store.save_checkpoint( + session_id=session_id, + checkpoint="search_done", + state={"search_result": search_result}, + ) + summary = search_result.get("summary") or {} yield StreamEvent( type="search_done", @@ -184,22 +331,47 @@ async def _dailypaper_stream(req: DailyPaperRequest): "items_count": len(search_result.get("items") or []), "queries_count": len(search_result.get("queries") or []), "unique_items": int(summary.get("unique_items") or 0), + "session_id": session_id, }, ) # Phase 2 — Build Report - yield StreamEvent(type="progress", data={"phase": "build", "message": "Building report..."}) - report = build_daily_paper_report(search_result=search_result, title=req.title, top_n=req.top_n) + if req.resume and isinstance(session_state.get("report"), dict): + report = dict(session_state.get("report") or {}) + yield StreamEvent( + type="progress", + data={"phase": "build", "message": "Resumed report from checkpoint"}, + ) + else: + yield StreamEvent(type="progress", data={"phase": "build", "message": "Building report..."}) + report = build_daily_paper_report( + search_result=search_result, title=req.title, top_n=req.top_n + ) + _pipeline_session_store.save_checkpoint( + session_id=session_id, + checkpoint="report_built", + state={"search_result": search_result, "report": report}, + ) + yield StreamEvent( type="report_built", data={ "queries_count": len(report.get("queries") or []), "global_top_count": len(report.get("global_top") or []), "report": report, + "session_id": session_id, }, ) - # Phase 3 — LLM Enrichment + query_items: List[Dict[str, Any]] = [] + paper_query_map: Dict[int, str] = {} + for query in report.get("queries") or []: + query_name = query.get("normalized_query") or query.get("raw_query") or "" + for item in query.get("top_items") or []: + query_items.append(item) + paper_query_map[id(item)] = query_name + + # Phase 3 — LLM Enrichment (pipeline) if req.enable_llm_analysis: features = normalize_llm_features(req.llm_features) if features: @@ -211,52 +383,42 @@ async def _dailypaper_stream(req: DailyPaperRequest): "daily_insight": "", } - summary_done = 0 - summary_total = 0 + llm_targets: set[int] = set() if "summary" in features or "relevance" in features: for query in report.get("queries") or []: - summary_total += len((query.get("top_items") or [])[:3]) + for item in (query.get("top_items") or [])[:3]: + llm_targets.add(id(item)) yield StreamEvent( type="progress", data={ "phase": "llm", "message": "Starting LLM enrichment...", - "total": summary_total, + "total": len(llm_targets), }, ) - for query in report.get("queries") or []: - query_name = query.get("normalized_query") or query.get("raw_query") or "" - top_items = (query.get("top_items") or [])[:3] - - if "summary" in features: - for item in top_items: - item["ai_summary"] = llm_service.summarize_paper( - title=item.get("title") or "", - abstract=item.get("snippet") or item.get("abstract") or "", - ) - summary_done += 1 - yield StreamEvent( - type="llm_summary", - data={ - "title": item.get("title") or "Untitled", - "query": query_name, - "ai_summary": item["ai_summary"], - "done": summary_done, - "total": summary_total, - }, - ) - - if "relevance" in features: - for item in top_items: - item["relevance"] = llm_service.assess_relevance( - paper=item, query=query_name - ) - if "summary" not in features: - summary_done += 1 - - if "trends" in features and top_items: + if llm_targets: + pipeline = EnrichmentPipeline( + steps=[LLMEnrichmentStep(llm_service=llm_service, features=features)] + ) + await pipeline.run( + query_items, + context=EnrichmentContext( + query="; ".join(cleaned_queries), + extra={ + "llm_target_ids": llm_targets, + "query_for_relevance": "; ".join(cleaned_queries), + }, + ), + ) + + if "trends" in features: + for query in report.get("queries") or []: + query_name = query.get("normalized_query") or query.get("raw_query") or "" + top_items = (query.get("top_items") or [])[:3] + if not top_items: + continue trend_text = llm_service.analyze_trends(topic=query_name, papers=top_items) llm_block["query_trends"].append({"query": query_name, "analysis": trend_text}) yield StreamEvent( @@ -278,6 +440,11 @@ async def _dailypaper_stream(req: DailyPaperRequest): yield StreamEvent(type="insight", data={"analysis": llm_block["daily_insight"]}) report["llm_analysis"] = llm_block + summary_done = sum( + 1 + for item in query_items + if id(item) in llm_targets and (item.get("ai_summary") or item.get("relevance")) + ) yield StreamEvent( type="llm_done", data={ @@ -286,7 +453,7 @@ async def _dailypaper_stream(req: DailyPaperRequest): }, ) - # Phase 4 — Judge + # Phase 4 — Judge + Filter (pipeline) if req.enable_judge: llm_service_j = get_llm_service() judge = PaperJudge(llm_service=llm_service_j) @@ -297,12 +464,17 @@ async def _dailypaper_stream(req: DailyPaperRequest): token_budget=req.judge_token_budget, ) selected = list(selection.get("selected") or []) - recommendation_count: Dict[str, int] = { - "must_read": 0, - "worth_reading": 0, - "skim": 0, - "skip": 0, - } + judge_targets: set[int] = set() + queries = list(report.get("queries") or []) + for row in selected: + query_index = int(row.get("query_index") or 0) + item_index = int(row.get("item_index") or 0) + if query_index >= len(queries): + continue + top_items = list(queries[query_index].get("top_items") or []) + if item_index >= len(top_items): + continue + judge_targets.add(id(top_items[item_index])) yield StreamEvent( type="progress", @@ -314,47 +486,32 @@ async def _dailypaper_stream(req: DailyPaperRequest): }, ) - queries = list(report.get("queries") or []) - for idx, row in enumerate(selected, start=1): - query_index = int(row.get("query_index") or 0) - item_index = int(row.get("item_index") or 0) - - if query_index >= len(queries): - continue + if judge_targets: + judge_pipeline = EnrichmentPipeline( + steps=[JudgeStep(judge=judge, n_runs=max(1, int(req.judge_runs)))] + ) + await judge_pipeline.run( + query_items, + context=EnrichmentContext( + query="; ".join(cleaned_queries), + extra={"judge_target_ids": judge_targets, "paper_query_map": paper_query_map}, + ), + ) - query = queries[query_index] - query_name = query.get("normalized_query") or query.get("raw_query") or "" - top_items = list(query.get("top_items") or []) - if item_index >= len(top_items): + recommendation_count: Dict[str, int] = { + "must_read": 0, + "worth_reading": 0, + "skim": 0, + "skip": 0, + } + for item in query_items: + if id(item) not in judge_targets: continue - - item = top_items[item_index] - if req.judge_runs > 1: - judgment = judge.judge_with_calibration( - paper=item, - query=query_name, - n_runs=max(1, int(req.judge_runs)), - ) - else: - judgment = judge.judge_single(paper=item, query=query_name) - - j_payload = judgment.to_dict() - item["judge"] = j_payload - rec = j_payload.get("recommendation") + j_payload = item.get("judge") if isinstance(item.get("judge"), dict) else {} + rec = str(j_payload.get("recommendation") or "") if rec in recommendation_count: recommendation_count[rec] += 1 - yield StreamEvent( - type="judge", - data={ - "query": query_name, - "title": item.get("title") or "Untitled", - "judge": j_payload, - "done": idx, - "total": len(selected), - }, - ) - for query in report.get("queries") or []: top_items = list(query.get("top_items") or []) if not top_items: @@ -375,12 +532,15 @@ async def _dailypaper_stream(req: DailyPaperRequest): } yield StreamEvent(type="judge_done", data=report["judge"]) - # Phase 4b — Filter: remove papers below "worth_reading" KEEP_RECOMMENDATIONS = {"must_read", "worth_reading"} yield StreamEvent( type="progress", data={"phase": "filter", "message": "Filtering papers by judge recommendation..."}, ) + + filter_pipeline = EnrichmentPipeline(steps=[FilterStep(keep=KEEP_RECOMMENDATIONS)]) + await filter_pipeline.run(query_items, context=EnrichmentContext()) + filter_log: List[Dict[str, Any]] = [] total_before = 0 total_after = 0 @@ -389,37 +549,40 @@ async def _dailypaper_stream(req: DailyPaperRequest): items_before = list(query.get("top_items") or []) total_before += len(items_before) kept: List[Dict[str, Any]] = [] - removed: List[Dict[str, Any]] = [] for item in items_before: - j = item.get("judge") - if isinstance(j, dict): - rec = j.get("recommendation", "") - if rec in KEEP_RECOMMENDATIONS: - kept.append(item) - else: - removed.append(item) - filter_log.append( - { - "query": query_name, - "title": item.get("title") or "Untitled", - "recommendation": rec, - "overall": j.get("overall"), - "action": "removed", - } - ) - else: - # No judge score — keep by default (unjudged papers) - kept.append(item) + if item.get("_filtered_out"): + j = item.get("judge") if isinstance(item.get("judge"), dict) else {} + filter_log.append( + { + "query": query_name, + "title": item.get("title") or "Untitled", + "recommendation": j.get("recommendation"), + "overall": j.get("overall"), + "action": "removed", + } + ) + continue + kept.append(item) total_after += len(kept) query["top_items"] = kept - # Also filter global_top + judge_by_key: Dict[str, Dict[str, Any]] = {} + for item in query_items: + if not isinstance(item.get("judge"), dict): + continue + key = f"{(item.get('url') or '').strip()}|{(item.get('title') or '').strip().lower()}" + if key: + judge_by_key[key] = item["judge"] + global_before = list(report.get("global_top") or []) global_kept = [] for item in global_before: + key = f"{(item.get('url') or '').strip()}|{(item.get('title') or '').strip().lower()}" + if key in judge_by_key: + item["judge"] = judge_by_key[key] j = item.get("judge") if isinstance(j, dict): - rec = j.get("recommendation", "") + rec = str(j.get("recommendation") or "") if rec in KEEP_RECOMMENDATIONS: global_kept.append(item) else: @@ -444,6 +607,42 @@ async def _dailypaper_stream(req: DailyPaperRequest): }, ) + _pipeline_session_store.save_checkpoint( + session_id=session_id, + checkpoint="enriched", + state={"search_result": search_result, "report": report}, + ) + + if req.require_approval: + preview_markdown = render_daily_paper_markdown(report) + pending_payload = { + "report": report, + "markdown": preview_markdown, + "markdown_path": None, + "json_path": None, + "notify_result": None, + "session_id": session_id, + "resumed": False, + "approval_status": "pending_approval", + } + _pipeline_session_store.update_status( + session_id=session_id, + status="pending_approval", + checkpoint="approval_pending", + state_patch={"search_result": search_result, "report": report}, + result=pending_payload, + ) + yield StreamEvent( + type="approval_required", + data={ + "phase": "approval", + "session_id": session_id, + "status": "pending_approval", + }, + ) + yield StreamEvent(type="result", data=pending_payload) + return + # Phase 5 — Persist + Notify yield StreamEvent(type="progress", data={"phase": "save", "message": "Saving to registry..."}) try: @@ -490,17 +689,22 @@ async def _dailypaper_stream(req: DailyPaperRequest): email_to_override=_validate_email_list(req.notify_email_to) or None, ) - yield StreamEvent( - type="result", - data={ - "report": report, - "markdown": markdown, - "markdown_path": markdown_path, - "json_path": json_path, - "notify_result": notify_result, - }, + result_payload = { + "report": report, + "markdown": markdown, + "markdown_path": markdown_path, + "json_path": json_path, + "notify_result": notify_result, + "session_id": session_id, + "resumed": False, + "approval_status": "approved", + } + _pipeline_session_store.save_result( + session_id=session_id, result=result_payload, status="completed" ) + yield StreamEvent(type="result", data=result_payload) + @router.post("/research/paperscool/daily") async def generate_daily_report(req: DailyPaperRequest): @@ -508,9 +712,15 @@ async def generate_daily_report(req: DailyPaperRequest): if not cleaned_queries: raise HTTPException(status_code=400, detail="queries is required") - # Fast sync path when no LLM/Judge — avoids SSE overhead - if not req.enable_llm_analysis and not req.enable_judge: - return _sync_daily_report(req, cleaned_queries) + # Fast sync path when no long-running step is requested — avoids SSE overhead + if ( + not req.enable_llm_analysis + and not req.enable_judge + and not req.require_approval + and not req.resume + and not req.session_id + ): + return await _sync_daily_report(req, cleaned_queries) # SSE streaming path for long-running operations return StreamingResponse( @@ -520,12 +730,163 @@ async def generate_daily_report(req: DailyPaperRequest): ) -def _sync_daily_report(req: DailyPaperRequest, cleaned_queries: List[str]): +@router.get("/research/paperscool/sessions/{session_id}", response_model=PipelineSessionResponse) +async def get_daily_session(session_id: str): + session = _pipeline_session_store.get_session(session_id) + if not session: + raise HTTPException(status_code=404, detail="session not found") + return PipelineSessionResponse(session=session) + + +@router.get("/research/paperscool/approvals", response_model=ApprovalQueueResponse) +async def list_pending_approvals(limit: int = 20): + rows = _pipeline_session_store.list_sessions( + workflow="paperscool_daily", + status="pending_approval", + limit=max(1, min(int(limit), 200)), + ) + items: List[Dict[str, Any]] = [] + for row in rows: + result = row.get("result") if isinstance(row.get("result"), dict) else {} + report = result.get("report") if isinstance(result.get("report"), dict) else {} + stats = report.get("stats") if isinstance(report.get("stats"), dict) else {} + items.append( + { + "session_id": row.get("session_id"), + "status": row.get("status"), + "checkpoint": row.get("checkpoint"), + "updated_at": row.get("updated_at"), + "title": report.get("title") or "DailyPaper Digest", + "query_count": int(stats.get("query_count") or 0), + "unique_items": int(stats.get("unique_items") or 0), + } + ) + return ApprovalQueueResponse(items=items) + + +def _finalize_approved_session(session: Dict[str, Any]) -> Dict[str, Any]: + payload = session.get("payload") if isinstance(session.get("payload"), dict) else {} + state = session.get("state") if isinstance(session.get("state"), dict) else {} + result = session.get("result") if isinstance(session.get("result"), dict) else {} + + report = state.get("report") if isinstance(state.get("report"), dict) else {} + if not report: + report = result.get("report") if isinstance(result.get("report"), dict) else {} + if not report: + raise HTTPException(status_code=400, detail="session has no report to approve") + + try: + report["registry_ingest"] = ingest_daily_report_to_registry(report) + except Exception as exc: + report["registry_ingest"] = {"error": str(exc)} + + if bool(payload.get("enable_judge")): + try: + report["judge_registry_ingest"] = persist_judge_scores_to_registry(report) + except Exception as exc: + report["judge_registry_ingest"] = {"error": str(exc)} + + _enqueue_repo_enrichment_async(report) + + markdown = render_daily_paper_markdown(report) + markdown_path = None + json_path = None + notify_result: Optional[Dict[str, Any]] = None + + if bool(payload.get("save")): + reporter = DailyPaperReporter( + output_dir=_sanitize_output_dir( + str(payload.get("output_dir") or "./reports/dailypaper") + ) + ) + artifacts = reporter.write( + report=report, + markdown=markdown, + formats=normalize_output_formats(payload.get("formats") or ["both"]), + slug=payload.get("title") or report.get("title") or "DailyPaper Digest", + ) + markdown_path = artifacts.markdown_path + json_path = artifacts.json_path + + if bool(payload.get("notify")): + notify_service = DailyPushService.from_env() + notify_result = notify_service.push_dailypaper( + report=report, + markdown=markdown, + markdown_path=markdown_path, + json_path=json_path, + channels_override=payload.get("notify_channels") or None, + email_to_override=_validate_email_list(payload.get("notify_email_to") or []) or None, + ) + + return { + "report": report, + "markdown": markdown, + "markdown_path": markdown_path, + "json_path": json_path, + "notify_result": notify_result, + "session_id": session.get("session_id"), + "resumed": False, + "approval_status": "approved", + } + + +@router.post( + "/research/paperscool/sessions/{session_id}/approve", response_model=PipelineSessionResponse +) +async def approve_daily_session(session_id: str): + session = _pipeline_session_store.get_session(session_id) + if not session: + raise HTTPException(status_code=404, detail="session not found") + if session.get("status") != "pending_approval": + raise HTTPException(status_code=409, detail="session is not pending approval") + + final_payload = _finalize_approved_session(session) + _pipeline_session_store.update_status( + session_id=session_id, + status="completed", + checkpoint="result", + state_patch={"approved_at": True}, + result=final_payload, + ) + updated = _pipeline_session_store.get_session(session_id) + return PipelineSessionResponse(session=updated or {}) + + +@router.post( + "/research/paperscool/sessions/{session_id}/reject", response_model=PipelineSessionResponse +) +async def reject_daily_session(session_id: str, req: ApprovalDecisionRequest): + session = _pipeline_session_store.get_session(session_id) + if not session: + raise HTTPException(status_code=404, detail="session not found") + if session.get("status") != "pending_approval": + raise HTTPException(status_code=409, detail="session is not pending approval") + + current_result = session.get("result") if isinstance(session.get("result"), dict) else {} + rejected_result = { + **current_result, + "session_id": session_id, + "approval_status": "rejected", + "rejected_reason": req.reason or "", + } + + _pipeline_session_store.update_status( + session_id=session_id, + status="rejected", + checkpoint="approval_rejected", + state_patch={"reject_reason": req.reason or ""}, + result=rejected_result, + ) + updated = _pipeline_session_store.get_session(session_id) + return PipelineSessionResponse(session=updated or {}) + + +async def _sync_daily_report(req: DailyPaperRequest, cleaned_queries: List[str]): """Original synchronous path for fast requests (no LLM/Judge).""" - workflow = PapersCoolTopicSearchWorkflow() effective_top_k = max(int(req.top_k_per_query), int(req.top_n), 1) try: - search_result = workflow.run( + search_result = await _run_topic_search( queries=cleaned_queries, sources=req.sources, branches=req.branches, diff --git a/src/paperbot/api/routes/research.py b/src/paperbot/api/routes/research.py index 5159d469..786a616f 100644 --- a/src/paperbot/api/routes/research.py +++ b/src/paperbot/api/routes/research.py @@ -3,7 +3,7 @@ import os import re from collections import Counter -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone from typing import Any, Dict, List, Optional, Tuple from fastapi import APIRouter, BackgroundTasks, HTTPException, Query @@ -27,6 +27,66 @@ _track_router = TrackRouter(research_store=_research_store, memory_store=_memory_store) _metric_collector: Optional[MemoryMetricCollector] = None _paper_store: Optional["PaperStore"] = None +_paper_search_service: Optional["PaperSearchService"] = None + +_DEADLINE_RADAR_DATA: List[Dict[str, Any]] = [ + { + "name": "KDD 2026", + "ccf_level": "A", + "field": "Data Mining", + "deadline": "2026-03-05T23:59:59+00:00", + "url": "https://kdd.org/kdd2026/", + "keywords": ["data mining", "recommendation", "graph mining"], + }, + { + "name": "ACL 2026", + "ccf_level": "A", + "field": "NLP", + "deadline": "2026-03-15T23:59:59+00:00", + "url": "https://2026.aclweb.org/", + "keywords": ["nlp", "llm", "language model", "retrieval"], + }, + { + "name": "CVPR 2026", + "ccf_level": "A", + "field": "Computer Vision", + "deadline": "2026-03-20T23:59:59+00:00", + "url": "https://cvpr.thecvf.com/", + "keywords": ["computer vision", "diffusion", "multimodal"], + }, + { + "name": "USENIX Security 2026", + "ccf_level": "A", + "field": "Security", + "deadline": "2026-03-28T23:59:59+00:00", + "url": "https://www.usenix.org/conference/usenixsecurity26", + "keywords": ["security", "privacy", "llm safety"], + }, + { + "name": "EMNLP 2026", + "ccf_level": "B", + "field": "NLP", + "deadline": "2026-05-10T23:59:59+00:00", + "url": "https://2026.emnlp.org/", + "keywords": ["nlp", "alignment", "reasoning"], + }, + { + "name": "NeurIPS 2026", + "ccf_level": "A", + "field": "Machine Learning", + "deadline": "2026-05-15T23:59:59+00:00", + "url": "https://neurips.cc/", + "keywords": ["machine learning", "llm", "optimization"], + }, + { + "name": "AAAI 2027", + "ccf_level": "A", + "field": "Artificial Intelligence", + "deadline": "2026-08-10T23:59:59+00:00", + "url": "https://aaai.org/conference/aaai/", + "keywords": ["ai", "agent", "reasoning"], + }, +] def _get_metric_collector() -> MemoryMetricCollector: @@ -47,6 +107,20 @@ def _get_paper_store() -> "PaperStore": return _paper_store +def _get_paper_search_service() -> "PaperSearchService": + """Lazy initialization of unified paper search service.""" + from paperbot.application.services.paper_search_service import PaperSearchService + from paperbot.infrastructure.adapters import build_adapter_registry + + global _paper_search_service + if _paper_search_service is None: + _paper_search_service = PaperSearchService( + adapters=build_adapter_registry(), + registry=_get_paper_store(), + ) + return _paper_search_service + + def _schedule_embedding_precompute( background_tasks: Optional[BackgroundTasks], *, @@ -125,6 +199,102 @@ def list_tracks( return TrackListResponse(user_id=user_id, tracks=tracks) +class DeadlineRadarResponse(BaseModel): + user_id: str + generated_at: str + items: List[Dict[str, Any]] + + +@router.get("/research/deadlines/radar", response_model=DeadlineRadarResponse) +def get_deadline_radar( + user_id: str = "default", + days: int = Query(180, ge=7, le=365), + ccf_levels: str = Query("A,B,C"), + field: Optional[str] = None, + limit: int = Query(20, ge=1, le=100), +): + levels = { + token.strip().upper() + for token in str(ccf_levels or "").split(",") + if token.strip().upper() in {"A", "B", "C"} + } + if not levels: + levels = {"A", "B", "C"} + + field_filter = str(field or "").strip().lower() + now = datetime.now(timezone.utc) + cutoff = now + timedelta(days=int(days)) + + tracks = _research_store.list_tracks(user_id=user_id, include_archived=False, limit=200) + track_tokens: Dict[int, set[str]] = {} + for track in tracks: + track_id = int(track.get("id") or 0) + if track_id <= 0: + continue + tokens = { + str(term).strip().lower() for term in (track.get("keywords") or []) if str(term).strip() + } + track_tokens[track_id] = tokens + + rows: List[Dict[str, Any]] = [] + for item in _DEADLINE_RADAR_DATA: + try: + deadline = datetime.fromisoformat(str(item.get("deadline") or "")) + if deadline.tzinfo is None: + deadline = deadline.replace(tzinfo=timezone.utc) + except Exception: + continue + + if deadline < now or deadline > cutoff: + continue + if str(item.get("ccf_level") or "").strip().upper() not in levels: + continue + if field_filter and field_filter not in str(item.get("field") or "").strip().lower(): + continue + + conf_keywords = { + str(k).strip().lower() for k in (item.get("keywords") or []) if str(k).strip() + } + + matched_tracks: List[Dict[str, Any]] = [] + for track in tracks: + track_id = int(track.get("id") or 0) + if track_id <= 0: + continue + overlap = sorted(conf_keywords & track_tokens.get(track_id, set())) + if overlap: + matched_tracks.append( + { + "track_id": track_id, + "track_name": str(track.get("name") or ""), + "matched_keywords": overlap, + } + ) + + workflow_query = ", ".join(item.get("keywords") or []) + days_left = max(0, int((deadline - now).total_seconds() // 86400)) + rows.append( + { + "name": str(item.get("name") or ""), + "ccf_level": str(item.get("ccf_level") or ""), + "field": str(item.get("field") or ""), + "deadline": deadline.isoformat(), + "days_left": days_left, + "url": str(item.get("url") or ""), + "keywords": sorted(conf_keywords), + "workflow_query": workflow_query, + "matched_tracks": matched_tracks, + } + ) + + rows.sort(key=lambda row: (int(row.get("days_left") or 0), str(row.get("name") or ""))) + return DeadlineRadarResponse( + user_id=user_id, + generated_at=now.isoformat(), + items=rows[: max(1, int(limit))], + ) + + @router.get("/research/tracks/active", response_model=TrackResponse) def get_active_track(user_id: str = "default"): track = _research_store.get_active_track(user_id=user_id) @@ -805,6 +975,15 @@ class SavedPapersResponse(BaseModel): items: List[Dict[str, Any]] +class TrackFeedResponse(BaseModel): + user_id: str + track_id: int + total: int + limit: int + offset: int + items: List[Dict[str, Any]] + + class PaperDetailResponse(BaseModel): detail: Dict[str, Any] @@ -831,13 +1010,46 @@ 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", + track_id: Optional[int] = None, sort_by: str = Query("saved_at"), limit: int = Query(200, ge=1, le=1000), ): - items = _research_store.list_saved_papers(user_id=user_id, sort_by=sort_by, limit=limit) + items = _research_store.list_saved_papers( + user_id=user_id, + track_id=track_id, + sort_by=sort_by, + limit=limit, + ) return SavedPapersResponse(user_id=user_id, items=items) +@router.get("/research/tracks/{track_id}/feed", response_model=TrackFeedResponse) +def get_track_feed( + track_id: int, + user_id: str = "default", + limit: int = Query(20, ge=1, le=100), + offset: int = Query(0, ge=0), +): + 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") + + payload = _research_store.list_track_feed( + user_id=user_id, + track_id=track_id, + limit=limit, + offset=offset, + ) + return TrackFeedResponse( + user_id=user_id, + track_id=track_id, + total=int(payload.get("total") or 0), + limit=limit, + offset=offset, + items=payload.get("items") or [], + ) + + @router.get("/research/papers/{paper_id}", response_model=PaperDetailResponse) def get_paper_detail(paper_id: str, user_id: str = "default"): detail = _research_store.get_paper_detail(paper_id=paper_id, user_id=user_id) @@ -881,6 +1093,7 @@ class ContextRequest(BaseModel): activate_track_id: Optional[int] = None # confirm switch: activates then uses it memory_limit: int = Field(8, ge=1, le=50) paper_limit: int = Field(8, ge=0, le=50) + sources: Optional[List[str]] = None offline: bool = False include_cross_track: bool = False stage: str = "auto" # auto/survey/writing/rebuttal @@ -907,13 +1120,26 @@ async def build_context(req: ContextRequest): raise HTTPException(status_code=404, detail="Track not found") Logger.info("Initializing context engine", file=LogFiles.HARVEST) + search_service = None + if not req.offline and req.paper_limit > 0: + try: + search_service = _get_paper_search_service() + except Exception as exc: + Logger.warning( + f"Failed to initialize PaperSearchService, fallback to legacy S2 path: {exc}", + file=LogFiles.HARVEST, + ) + engine = ContextEngine( research_store=_research_store, memory_store=_memory_store, + paper_store=_get_paper_store(), + search_service=search_service, track_router=_track_router, config=ContextEngineConfig( memory_limit=req.memory_limit, paper_limit=req.paper_limit, + search_sources=req.sources, offline=req.offline, stage=req.stage, exploration_ratio=( diff --git a/src/paperbot/api/streaming.py b/src/paperbot/api/streaming.py index 2c870637..de694cc6 100644 --- a/src/paperbot/api/streaming.py +++ b/src/paperbot/api/streaming.py @@ -15,10 +15,20 @@ import json from dataclasses import dataclass from datetime import datetime, timezone +from enum import Enum from typing import Any, AsyncGenerator, Dict, Optional from uuid import uuid4 +class StandardEvent(str, Enum): + STATUS = "status" + PROGRESS = "progress" + TOOL = "tool" + RESULT = "result" + ERROR = "error" + DONE = "done" + + def _new_stream_id(prefix: str) -> str: return f"{prefix}_{uuid4().hex[:12]}" @@ -28,6 +38,7 @@ class StreamEvent: """SSE event structure.""" type: str # progress, result, error, done + event: Optional[str] = None # canonical event kind for unified frontend handling data: Any = None message: Optional[str] = None envelope: Optional[Dict[str, Any]] = None @@ -36,6 +47,7 @@ def to_sse(self) -> str: """Convert to SSE format.""" payload = { "type": self.type, + "event": self.event, "data": self.data, "message": self.message, "envelope": self.envelope, @@ -48,6 +60,46 @@ def sse_done() -> str: return "data: [DONE]\n\n" +def _canonical_event_kind( + *, + event_type: str, + data: Any, + explicit_event: Optional[str], +) -> str: + if explicit_event: + return str(explicit_event) + + t = str(event_type or "").strip().lower() + if t in {"error", "failed", "failure"}: + return StandardEvent.ERROR.value + if t in {"result", "final", "final_result"}: + return StandardEvent.RESULT.value + if t in {"done", "completed", "complete"}: + return StandardEvent.DONE.value + if t == "status": + return StandardEvent.STATUS.value + if t.startswith("tool"): + return StandardEvent.TOOL.value + if t in { + "progress", + "search_done", + "report_built", + "llm_summary", + "llm_done", + "trend", + "insight", + "judge", + "judge_done", + "filter_done", + }: + return StandardEvent.PROGRESS.value + + if isinstance(data, dict): + if any(k in data for k in ("phase", "delta", "done", "total")): + return StandardEvent.PROGRESS.value + return StandardEvent.STATUS.value + + def _with_envelope( event: StreamEvent, *, @@ -56,7 +108,16 @@ def _with_envelope( trace_id: str, seq: int, ) -> StreamEvent: + canonical_event = _canonical_event_kind( + event_type=event.type, + data=event.data, + explicit_event=event.event, + ) + event.event = canonical_event + if event.envelope: + if "event" not in event.envelope: + event.envelope["event"] = canonical_event return event phase = None @@ -69,6 +130,7 @@ def _with_envelope( "trace_id": trace_id, "seq": seq, "phase": phase, + "event": canonical_event, "ts": datetime.now(timezone.utc).isoformat(), } return event diff --git a/src/paperbot/application/services/enrichment_pipeline.py b/src/paperbot/application/services/enrichment_pipeline.py index e0fd09e6..ed7667e2 100644 --- a/src/paperbot/application/services/enrichment_pipeline.py +++ b/src/paperbot/application/services/enrichment_pipeline.py @@ -1,7 +1,7 @@ """EnrichmentPipeline — Chain of Responsibility for paper enrichment. Each step processes a paper and can add judge scores, summaries, -repo discovery results, etc. +or post-filter flags. """ from __future__ import annotations @@ -51,4 +51,75 @@ async def run( except Exception as e: title = str(paper.get("title", ""))[:60] step_name = type(step).__name__ - logger.warning(f"Enrichment step {step_name} failed for '{title}': {e}") + logger.warning(f"Enrichment step {step_name} failed for {title}: {e}") + + +class LLMEnrichmentStep: + """Attach LLM summary/relevance features to selected papers.""" + + def __init__(self, *, llm_service=None, features: Optional[List[str]] = None): + from paperbot.application.services.llm_service import get_llm_service + + self._llm = llm_service or get_llm_service() + self._features = set(features or ["summary"]) + + async def process(self, paper: Dict[str, Any], context: EnrichmentContext) -> None: + target_ids = context.extra.get("llm_target_ids") + if isinstance(target_ids, set) and id(paper) not in target_ids: + return + + title = str(paper.get("title") or "") + abstract = str(paper.get("snippet") or paper.get("abstract") or "") + + if "summary" in self._features: + paper["ai_summary"] = self._llm.summarize_paper(title=title, abstract=abstract) + + if "relevance" in self._features: + query = str(context.extra.get("query_for_relevance") or context.query or "") + paper["relevance"] = self._llm.assess_relevance(paper=paper, query=query) + + +class JudgeStep: + """Attach judge scores to selected papers.""" + + def __init__(self, *, judge=None, n_runs: int = 1): + if judge is None: + from paperbot.application.services.llm_service import get_llm_service + from paperbot.application.workflows.analysis.paper_judge import PaperJudge + + judge = PaperJudge(llm_service=get_llm_service()) + self._judge = judge + self._n_runs = max(1, int(n_runs)) + + async def process(self, paper: Dict[str, Any], context: EnrichmentContext) -> None: + target_ids = context.extra.get("judge_target_ids") + if isinstance(target_ids, set) and id(paper) not in target_ids: + return + + query_map = context.extra.get("paper_query_map") or {} + query = str(query_map.get(id(paper)) or context.query or "") + + if self._n_runs > 1: + judgment = self._judge.judge_with_calibration( + paper=paper, + query=query, + n_runs=self._n_runs, + ) + else: + judgment = self._judge.judge_single(paper=paper, query=query) + paper["judge"] = judgment.to_dict() + + +class FilterStep: + """Mark papers as filtered when recommendation is not in keep-set.""" + + def __init__(self, keep: Optional[set[str]] = None): + self._keep = keep or {"must_read", "worth_reading"} + + async def process(self, paper: Dict[str, Any], context: EnrichmentContext) -> None: + judge = paper.get("judge") + if not isinstance(judge, dict): + return + rec = str(judge.get("recommendation") or "").strip().lower() + if rec and rec not in self._keep: + paper["_filtered_out"] = True diff --git a/src/paperbot/application/services/identity_resolver.py b/src/paperbot/application/services/identity_resolver.py index 2a5d6c81..f3e7f7e9 100644 --- a/src/paperbot/application/services/identity_resolver.py +++ b/src/paperbot/application/services/identity_resolver.py @@ -105,7 +105,6 @@ def resolve( select(PaperModel).where( or_( PaperModel.url.in_(url_candidates), - PaperModel.external_url.in_(url_candidates), PaperModel.pdf_url.in_(url_candidates), ) ) diff --git a/src/paperbot/application/services/llm_service.py b/src/paperbot/application/services/llm_service.py index e2f34cd1..71df5419 100644 --- a/src/paperbot/application/services/llm_service.py +++ b/src/paperbot/application/services/llm_service.py @@ -6,7 +6,12 @@ from typing import Any, Dict, Generator, List, Optional, Sequence from paperbot.application.prompts import PromptRegistry +from paperbot.application.services.provider_resolver import ( + ProviderResolver, + RouterBackedProviderResolver, +) from paperbot.infrastructure.llm.router import ModelRouter +from paperbot.infrastructure.stores.llm_usage_store import LLMUsageStore logger = logging.getLogger(__name__) @@ -17,16 +22,25 @@ class LLMService: def __init__( self, router: Optional[ModelRouter] = None, + provider_resolver: Optional[ProviderResolver] = None, prompt_registry: Optional[PromptRegistry] = None, + usage_store: Optional[LLMUsageStore] = None, *, enable_cache: bool = True, raise_errors: bool = False, ) -> None: - self._router = router or ModelRouter.from_env() + resolved_router = router or ModelRouter.from_env() + self._provider_resolver = provider_resolver or RouterBackedProviderResolver(resolved_router) self._prompts = prompt_registry or PromptRegistry() self._enable_cache = enable_cache self._raise_errors = raise_errors self._cache: Dict[str, str] = {} + self._usage_store = usage_store + if self._usage_store is None: + try: + self._usage_store = LLMUsageStore() + except Exception: + self._usage_store = None def complete( self, @@ -42,8 +56,15 @@ def complete( return self._cache[cache_key] try: - provider = self._router.get_provider(task_type) + provider = self._provider_resolver.get_provider(task_type) result = (provider.invoke_simple(system, user, **kwargs) or "").strip() + self._record_usage( + task_type=task_type, + provider=provider, + system=system, + user=user, + completion=result, + ) except Exception as exc: # pragma: no cover - exercised via fallback tests logger.warning("LLM complete failed task_type=%s error=%s", task_type, exc) if self._raise_errors: @@ -63,12 +84,24 @@ def stream( **kwargs, ) -> Generator[str, None, None]: try: - provider = self._router.get_provider(task_type) + provider = self._provider_resolver.get_provider(task_type) messages = [ {"role": "system", "content": system}, {"role": "user", "content": user}, ] - yield from provider.stream_invoke(messages, **kwargs) + chunks: List[str] = [] + for chunk in provider.stream_invoke(messages, **kwargs): + text = str(chunk or "") + if text: + chunks.append(text) + yield text + self._record_usage( + task_type=task_type, + provider=provider, + system=system, + user=user, + completion="".join(chunks), + ) except Exception as exc: # pragma: no cover - stream failures are runtime/network specific logger.warning("LLM stream failed task_type=%s error=%s", task_type, exc) if self._raise_errors: @@ -151,10 +184,53 @@ def _cache_key(self, *, task_type: str, system: str, user: str, kwargs: Dict[str ) return hashlib.sha256(payload.encode("utf-8")).hexdigest() + def _record_usage( + self, + *, + task_type: str, + provider: Any, + system: str, + user: str, + completion: str, + ) -> None: + if self._usage_store is None: + return + + prompt_tokens = _estimate_tokens(system) + _estimate_tokens(user) + completion_tokens = _estimate_tokens(completion) + + try: + info = provider.info + provider_name = str(getattr(info, "provider_name", "unknown") or "unknown") + model_name = str(getattr(info, "model_name", "") or "") + except Exception: + provider_name = "unknown" + model_name = "" + + cost_usd = _estimate_cost_usd( + provider_name=provider_name, + model_name=model_name, + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + ) + + try: + self._usage_store.record_usage( + task_type=task_type, + provider_name=provider_name, + model_name=model_name, + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + estimated_cost_usd=cost_usd, + metadata={"estimated": True}, + ) + except Exception: + return + def describe_task_provider(self, task_type: str = "default") -> Dict[str, Any]: """Expose selected provider metadata for auditing/judge reports.""" try: - provider = self._router.get_provider(task_type) + provider = self._provider_resolver.get_provider(task_type) info = provider.info return { "provider_name": info.provider_name, @@ -244,3 +320,34 @@ def _tokenize(text: str) -> List[str]: seen.add(token) dedup.append(token) return dedup + + +def _estimate_tokens(text: str) -> int: + raw = str(text or "") + if not raw: + return 0 + # Lightweight heuristic: ~4 chars/token for mixed English+code prompts. + return max(1, int(len(raw) / 4)) + + +def _estimate_cost_usd( + *, provider_name: str, model_name: str, prompt_tokens: int, completion_tokens: int +) -> float: + provider = (provider_name or "").lower() + model = (model_name or "").lower() + + in_price = 0.0 + out_price = 0.0 + + if provider == "openai" and "gpt-4o-mini" in model: + in_price, out_price = 0.15, 0.60 + elif provider == "openai" and "gpt-4o" in model: + in_price, out_price = 2.50, 10.00 + elif provider == "anthropic" and "claude-3-5-sonnet" in model: + in_price, out_price = 3.00, 15.00 + elif provider == "deepseek": + in_price, out_price = 0.55, 2.19 + elif provider == "ollama": + in_price, out_price = 0.0, 0.0 + + return ((prompt_tokens * in_price) + (completion_tokens * out_price)) / 1_000_000 diff --git a/src/paperbot/application/services/provider_resolver.py b/src/paperbot/application/services/provider_resolver.py new file mode 100644 index 00000000..d8e3bb17 --- /dev/null +++ b/src/paperbot/application/services/provider_resolver.py @@ -0,0 +1,22 @@ +from __future__ import annotations + +from typing import Protocol + +from paperbot.infrastructure.llm.providers.base import LLMProvider +from paperbot.infrastructure.llm.router import ModelRouter + + +class ProviderResolver(Protocol): + """Resolve task_type -> provider instance.""" + + def get_provider(self, task_type: str = "default") -> LLMProvider: ... + + +class RouterBackedProviderResolver: + """Adapter that keeps existing ModelRouter behavior behind a resolver interface.""" + + def __init__(self, router: ModelRouter): + self._router = router + + def get_provider(self, task_type: str = "default") -> LLMProvider: + return self._router.get_provider(task_type) diff --git a/src/paperbot/application/workflows/unified_topic_search.py b/src/paperbot/application/workflows/unified_topic_search.py new file mode 100644 index 00000000..bd5fb8f6 --- /dev/null +++ b/src/paperbot/application/workflows/unified_topic_search.py @@ -0,0 +1,356 @@ +from __future__ import annotations + +import asyncio +import math +import re +from datetime import datetime, timezone +from typing import Any, Dict, Iterable, List, Optional, Sequence + +from paperbot.application.services.paper_search_service import PaperSearchService, SearchResult + + +_QUERY_ALIASES = { + "icl压缩": "icl compression", + "icl 压缩": "icl compression", + "icl 隐式偏置": "icl implicit bias", + "icl隐式偏置": "icl implicit bias", + "kv cache加速": "kv cache acceleration", + "kv cache 加速": "kv cache acceleration", +} + +_SOURCE_ALIASES = { + "papers_cool": "papers_cool", + "paperscool": "papers_cool", + "arxiv": "arxiv", + "arxiv_api": "arxiv", + "hf_daily": "hf_daily", + "openalex": "openalex", + "semantic_scholar": "semantic_scholar", + "s2": "semantic_scholar", +} + + +def make_default_search_service(*, registry=None) -> PaperSearchService: + from paperbot.infrastructure.adapters import build_adapter_registry + + return PaperSearchService(adapters=build_adapter_registry(), registry=registry) + + +def normalize_topic_sources(sources: Sequence[str] | None) -> List[str]: + names: List[str] = [] + seen: set[str] = set() + for raw in sources or ["papers_cool"]: + normalized = _SOURCE_ALIASES.get((raw or "").strip().lower()) + if not normalized or normalized in seen: + continue + seen.add(normalized) + names.append(normalized) + if not names: + names = ["papers_cool"] + return names + + +def _normalize_query(query: str) -> str: + base = re.sub(r"\s+", " ", (query or "").strip()).lower() + return _QUERY_ALIASES.get(base, base) + + +def _tokenize_query(query: str) -> List[str]: + seen: set[str] = set() + tokens: List[str] = [] + for token in re.findall(r"[a-z0-9]+", query.lower()): + if token in seen: + continue + seen.add(token) + tokens.append(token) + return tokens + + +def _unique_preserve(values: Iterable[str]) -> List[str]: + seen: set[str] = set() + out: List[str] = [] + for value in values: + v = (value or "").strip() + if not v or v in seen: + continue + seen.add(v) + out.append(v) + return out + + +def _matched_keywords(paper: Dict[str, Any], tokens: List[str]) -> List[str]: + if not tokens: + return [] + blob = " ".join( + [ + str(paper.get("title") or ""), + str(paper.get("abstract") or paper.get("snippet") or ""), + " ".join(str(v) for v in (paper.get("keywords") or [])), + " ".join(str(v) for v in (paper.get("fields_of_study") or [])), + ] + ).lower() + return [tok for tok in tokens if tok in blob] + + +def _score_paper(paper: Dict[str, Any], query_tokens: List[str], matched: List[str]) -> float: + token_score = (len(matched) / max(1, len(query_tokens))) * 3.0 + citation_count = 0 + try: + citation_count = int(paper.get("citation_count") or 0) + except Exception: + citation_count = 0 + citation_score = min(1.25, math.log10(max(1, citation_count + 1)) / 3.0) + + year = None + try: + year = int(paper.get("year")) if paper.get("year") is not None else None + except Exception: + year = None + if year is not None: + age = max(0, datetime.now(timezone.utc).year - year) + recency_score = max(0.0, 1.0 - age / 10.0) + else: + recency_score = 0.25 + + return round(token_score + citation_score + recency_score, 4) + + +def _paper_id_from_item(item: Dict[str, Any]) -> str: + if item.get("canonical_id"): + return str(item["canonical_id"]) + + identities = item.get("identities") or [] + if isinstance(identities, list): + preferred = ["semantic_scholar", "arxiv", "openalex", "papers_cool", "hf_daily", "doi"] + by_source: Dict[str, str] = {} + for ident in identities: + if not isinstance(ident, dict): + continue + source = str(ident.get("source") or "").strip().lower() + external_id = str(ident.get("external_id") or "").strip() + if source and external_id and source not in by_source: + by_source[source] = external_id + for source in preferred: + if source in by_source: + return by_source[source] + if by_source: + return next(iter(by_source.values())) + + return "" + + +def _search_result_to_items( + *, + search_result: SearchResult, + normalized_query: str, + query_tokens: List[str], + branches: Sequence[str], + fallback_sources: List[str], + min_score: float, +) -> List[Dict[str, Any]]: + rows: List[Dict[str, Any]] = [] + + for paper in search_result.papers: + p = paper.to_dict() + key = p.get("title_hash") or p.get("title") + provenances = search_result.provenance.get(str(key), fallback_sources) + matched = _matched_keywords(p, query_tokens) + score = _score_paper(p, query_tokens, matched) + if score < float(min_score or 0.0): + continue + + url = str(p.get("url") or "").strip() + venue = str(p.get("venue") or "").strip() + rows.append( + { + "paper_id": _paper_id_from_item(p), + "title": str(p.get("title") or "").strip(), + "url": url, + "external_url": url, + "pdf_url": str(p.get("pdf_url") or "").strip(), + "authors": list(p.get("authors") or []), + "subject_or_venue": venue, + "published_at": p.get("publication_date") or (str(p.get("year")) if p.get("year") else ""), + "snippet": str(p.get("abstract") or "").strip(), + "keywords": _unique_preserve( + [ + *[str(v) for v in (p.get("keywords") or [])], + *[str(v) for v in (p.get("fields_of_study") or [])], + ] + ), + "branches": list(branches or ["arxiv", "venue"]), + "sources": _unique_preserve([str(v) for v in provenances]) or fallback_sources, + "matched_keywords": matched, + "matched_queries": [normalized_query], + "score": score, + "pdf_stars": 0, + "kimi_stars": 0, + "alternative_urls": [], + } + ) + + rows.sort(key=lambda row: float(row.get("score") or 0.0), reverse=True) + return rows + + +def _merge_item(target: Dict[str, Any], incoming: Dict[str, Any]) -> None: + target["matched_queries"] = _unique_preserve( + [*target.get("matched_queries", []), *incoming.get("matched_queries", [])] + ) + target["matched_keywords"] = _unique_preserve( + [*target.get("matched_keywords", []), *incoming.get("matched_keywords", [])] + ) + target["branches"] = _unique_preserve([*target.get("branches", []), *incoming.get("branches", [])]) + target["sources"] = _unique_preserve([*target.get("sources", []), *incoming.get("sources", [])]) + target["keywords"] = _unique_preserve([*target.get("keywords", []), *incoming.get("keywords", [])]) + target["authors"] = _unique_preserve([*target.get("authors", []), *incoming.get("authors", [])]) + + incoming_url = str(incoming.get("url") or "").strip() + target_url = str(target.get("url") or "").strip() + if incoming_url and incoming_url != target_url: + target["alternative_urls"] = _unique_preserve( + [*target.get("alternative_urls", []), incoming_url] + ) + + if float(incoming.get("score") or 0.0) > float(target.get("score") or 0.0): + target["score"] = float(incoming.get("score") or 0.0) + + +async def run_unified_topic_search( + *, + queries: Sequence[str], + branches: Sequence[str] = ("arxiv", "venue"), + sources: Sequence[str] = ("papers_cool",), + top_k_per_query: int = 5, + show_per_branch: int = 25, + min_score: float = 0.0, + search_service: Optional[PaperSearchService] = None, + persist: bool = False, +) -> Dict[str, Any]: + normalized_sources = normalize_topic_sources(sources) + + query_specs: List[Dict[str, Any]] = [] + seen_queries: set[str] = set() + for raw in queries: + raw_query = (raw or "").strip() + if not raw_query: + continue + normalized_query = _normalize_query(raw_query) + if normalized_query in seen_queries: + continue + seen_queries.add(normalized_query) + query_specs.append( + { + "raw_query": raw_query, + "normalized_query": normalized_query, + "tokens": _tokenize_query(normalized_query), + } + ) + + if not query_specs: + return { + "source": "papers.cool", + "fetched_at": datetime.now(timezone.utc).isoformat(), + "sources": normalized_sources, + "queries": [], + "items": [], + "summary": { + "unique_items": 0, + "total_query_hits": 0, + "top_titles": [], + "query_highlights": [], + "source_breakdown": {}, + "source_errors": [], + }, + } + + service = search_service or make_default_search_service() + max_results = max(1, int(show_per_branch)) + + tasks = [ + service.search( + spec["normalized_query"], + sources=normalized_sources, + max_results=max_results, + persist=bool(persist), + ) + for spec in query_specs + ] + search_results = await asyncio.gather(*tasks) + + query_views: List[Dict[str, Any]] = [] + aggregated: List[Dict[str, Any]] = [] + by_key: Dict[str, Dict[str, Any]] = {} + + for spec, search_result in zip(query_specs, search_results): + query_items = _search_result_to_items( + search_result=search_result, + normalized_query=spec["normalized_query"], + query_tokens=spec["tokens"], + branches=branches, + fallback_sources=normalized_sources, + min_score=min_score, + ) + + query_views.append( + { + "raw_query": spec["raw_query"], + "normalized_query": spec["normalized_query"], + "tokens": spec["tokens"], + "total_hits": len(query_items), + "items": query_items[: max(0, int(top_k_per_query))], + } + ) + + for item in query_items: + url = str(item.get("url") or "").strip().lower() + title = str(item.get("title") or "").strip().lower() + key = url or title + if not key: + continue + existing = by_key.get(key) + if existing is None: + cloned = dict(item) + by_key[key] = cloned + aggregated.append(cloned) + else: + _merge_item(existing, item) + + aggregated.sort(key=lambda row: float(row.get("score") or 0.0), reverse=True) + + query_highlights: List[Dict[str, Any]] = [] + total_query_hits = 0 + for row in query_views: + total_hits = int(row.get("total_hits") or 0) + total_query_hits += total_hits + top_item = (row.get("items") or [None])[0] + query_highlights.append( + { + "raw_query": row.get("raw_query") or "", + "normalized_query": row.get("normalized_query") or "", + "hit_count": total_hits, + "top_title": (top_item or {}).get("title") or "", + "top_keywords": ((top_item or {}).get("matched_keywords") or [])[:5], + } + ) + + source_breakdown: Dict[str, int] = {} + for item in aggregated: + for source in item.get("sources") or []: + source_breakdown[source] = source_breakdown.get(source, 0) + 1 + + return { + "source": "papers.cool", + "fetched_at": datetime.now(timezone.utc).isoformat(), + "sources": normalized_sources, + "queries": query_views, + "items": aggregated, + "summary": { + "unique_items": len(aggregated), + "total_query_hits": total_query_hits, + "top_titles": [item.get("title") or "" for item in aggregated[:5]], + "query_highlights": query_highlights, + "source_breakdown": source_breakdown, + "source_errors": [], + }, + } diff --git a/src/paperbot/context_engine/engine.py b/src/paperbot/context_engine/engine.py index b642da32..bb67b4fa 100644 --- a/src/paperbot/context_engine/engine.py +++ b/src/paperbot/context_engine/engine.py @@ -347,6 +347,7 @@ class ContextEngineConfig: paper_limit: int = 8 offline: bool = False stage: str = "auto" + search_sources: Optional[List[str]] = None exploration_ratio: Optional[float] = None diversity_strength: Optional[float] = None track_router: TrackRouterConfig = field(default_factory=TrackRouterConfig) @@ -358,6 +359,7 @@ def __init__( *, research_store: Optional[SqlAlchemyResearchStore] = None, memory_store: Optional[SqlAlchemyMemoryStore] = None, + paper_store: Optional[Any] = None, paper_searcher: Optional[Any] = None, search_service: Optional[Any] = None, track_router: Optional[TrackRouter] = None, @@ -365,6 +367,7 @@ def __init__( ): self.research_store = research_store or SqlAlchemyResearchStore() self.memory_store = memory_store or SqlAlchemyMemoryStore() + self.paper_store = paper_store self.paper_searcher = paper_searcher self.search_service = search_service self.config = config or ContextEngineConfig() @@ -374,6 +377,39 @@ def __init__( config=self.config.track_router, ) + def _attach_latest_judge(self, papers: List[Dict[str, Any]]) -> None: + ids: List[int] = [] + for paper in papers: + pid = str(paper.get("paper_id") or "").strip() + if pid.isdigit(): + ids.append(int(pid)) + if not ids: + return + + if self.paper_store is None: + from paperbot.infrastructure.stores.paper_store import PaperStore + + self.paper_store = PaperStore(auto_create_schema=False) + + judge_map = self.paper_store.get_latest_judge_scores(ids) + for paper in papers: + pid = str(paper.get("paper_id") or "").strip() + if pid.isdigit() and int(pid) in judge_map: + paper["latest_judge"] = judge_map[int(pid)] + + @staticmethod + def _attach_feedback_flags( + papers: List[Dict[str, Any]], *, saved_ids: set[str], liked_ids: set[str] + ) -> None: + for paper in papers: + pid = str(paper.get("paper_id") or "").strip() + if not pid: + continue + if pid in saved_ids: + paper["is_saved"] = True + if pid in liked_ids: + paper["is_liked"] = True + async def build_context_pack( self, *, @@ -522,13 +558,19 @@ async def build_context_pack( # Prefer PaperSearchService if available if self.search_service is not None: + selected_sources = [ + str(x).strip() for x in (self.config.search_sources or []) if str(x).strip() + ] + if not selected_sources: + selected_sources = ["semantic_scholar"] + Logger.info( f"Using PaperSearchService for query='{merged_query}'", file=LogFiles.HARVEST, ) search_result = await self.search_service.search( merged_query, - sources=["semantic_scholar"], + sources=selected_sources, max_results=fetch_limit, persist=True, ) @@ -556,7 +598,9 @@ async def build_context_pack( f"Searching papers with query='{merged_query}', limit={fetch_limit}", file=LogFiles.HARVEST, ) - resp = await asyncio.to_thread(searcher.search_papers, merged_query, fetch_limit) + resp = await asyncio.to_thread( + searcher.search_papers, merged_query, fetch_limit + ) papers_count = len(getattr(resp, "papers", []) or []) Logger.info(f"Search returned {papers_count} papers", file=LogFiles.HARVEST) @@ -597,6 +641,20 @@ async def build_context_pack( seen_titles.add(tkey) filtered.append(p) + try: + self._attach_latest_judge(filtered) + except Exception as exc: + Logger.warning( + f"Failed to attach latest judge scores: {exc}", + file=LogFiles.HARVEST, + ) + + self._attach_feedback_flags( + filtered, + saved_ids=set(saved_ids), + liked_ids=set(liked_ids), + ) + for p in filtered: pid = str(p.get("paper_id") or "").strip() if not pid: @@ -627,6 +685,7 @@ async def build_context_pack( ) except Exception as e: import traceback + tb = traceback.format_exc() Logger.error(f"Error fetching papers: {e}\n{tb}", file=LogFiles.HARVEST) papers = [] @@ -678,4 +737,21 @@ async def build_context_pack( } async def close(self) -> None: + if self.search_service is not None: + close_fn = getattr(self.search_service, "close", None) + if callable(close_fn): + try: + maybe_coro = close_fn() + if asyncio.iscoroutine(maybe_coro): + await maybe_coro + except Exception: + pass + + if self.paper_store is not None: + close_fn = getattr(self.paper_store, "close", None) + if callable(close_fn): + try: + close_fn() + except Exception: + pass return None diff --git a/src/paperbot/infrastructure/llm/router.py b/src/paperbot/infrastructure/llm/router.py index f6fd03dc..0e460f80 100644 --- a/src/paperbot/infrastructure/llm/router.py +++ b/src/paperbot/infrastructure/llm/router.py @@ -15,7 +15,7 @@ import logging import os from dataclasses import dataclass, field -from typing import Dict, Optional, Any, Callable +from typing import Dict, Optional, Any from enum import Enum from .providers.base import LLMProvider @@ -54,6 +54,7 @@ class ModelConfig: model: str cost_tier: int = 1 api_key_env: str = "OPENAI_API_KEY" + api_key: Optional[str] = None base_url: Optional[str] = None max_tokens: int = 4096 @@ -150,7 +151,7 @@ def get_provider(self, task_type: str = "default") -> LLMProvider: def _create_provider(self, config: ModelConfig) -> LLMProvider: """根据配置创建 Provider""" - api_key = os.getenv(config.api_key_env, "") + api_key = str(config.api_key or "").strip() or os.getenv(config.api_key_env, "") if config.provider == "openai": from .providers.openai_provider import OpenAIProvider @@ -191,10 +192,77 @@ def list_models(self) -> Dict[str, str]: """列出所有配置的模型""" return {name: f"{cfg.provider}:{cfg.model}" for name, cfg in self.config.models.items()} + @staticmethod + def _normalize_provider(vendor: str) -> str: + text = str(vendor or "").strip().lower() + if text in {"openai", "openai_compatible", "openai-compatible"}: + return "openai" + if text in {"anthropic", "ollama"}: + return text + return "openai" + + @classmethod + def _from_model_registry(cls) -> Optional["ModelRouter"]: + """Build router from user-managed model_endpoints registry if available.""" + try: + from paperbot.infrastructure.stores.model_endpoint_store import ModelEndpointStore + + store = ModelEndpointStore(auto_create_schema=False) + rows = store.list_endpoints(enabled_only=True, include_secrets=True) + except Exception as exc: + logger.warning("Model registry unavailable, fallback to env routing: %s", exc) + return None + + if not rows: + return None + + models: Dict[str, ModelConfig] = {} + fallback_name: Optional[str] = None + + for row in rows: + name = str(row.get("name") or "").strip() + model_list = [str(x).strip() for x in (row.get("models") or []) if str(x).strip()] + if not name or not model_list: + continue + + models[name] = ModelConfig( + provider=cls._normalize_provider(str(row.get("vendor") or "openai_compatible")), + model=model_list[0], + cost_tier=1, + api_key_env=str(row.get("api_key_env") or "OPENAI_API_KEY"), + api_key=(str(row.get("api_key") or "").strip() or None), + base_url=(str(row.get("base_url") or "").strip() or None), + ) + if bool(row.get("is_default")) and fallback_name is None: + fallback_name = name + + if not models: + return None + + fallback_name = fallback_name or next(iter(models.keys())) + router = cls(RouterConfig(models=models, fallback_model=fallback_name)) + + allowed_tasks = {str(t.value) for t in TaskType} + for row in rows: + name = str(row.get("name") or "").strip() + if name not in models: + continue + for task in row.get("task_types") or []: + task_name = str(task).strip().lower() + if task_name in allowed_tasks: + router.set_task_routing(task_name, name) + + router.set_task_routing("default", fallback_name) + return router + @classmethod def from_env(cls) -> "ModelRouter": """从环境变量创建默认路由器,并兼容 OpenAI-compatible 自定义端点。""" + registry_router = cls._from_model_registry() + if registry_router is not None: + return registry_router + default_model_override = os.getenv("LLM_DEFAULT_MODEL") reasoning_model_override = os.getenv("LLM_REASONING_MODEL") diff --git a/src/paperbot/infrastructure/queue/arq_worker.py b/src/paperbot/infrastructure/queue/arq_worker.py index 7256b183..687e61ae 100644 --- a/src/paperbot/infrastructure/queue/arq_worker.py +++ b/src/paperbot/infrastructure/queue/arq_worker.py @@ -389,16 +389,21 @@ async def daily_papers_job( render_daily_paper_markdown, ) from paperbot.application.services.daily_push_service import DailyPushService - from paperbot.application.workflows.paperscool_topic_search import PapersCoolTopicSearchWorkflow + from paperbot.application.workflows.unified_topic_search import ( + make_default_search_service, + run_unified_topic_search, + ) from paperbot.workflows.feed import ScholarFeedService - search_workflow = PapersCoolTopicSearchWorkflow() - search_result = search_workflow.run( + search_service = make_default_search_service() + search_result = await run_unified_topic_search( queries=job_queries, sources=job_sources, branches=job_branches, top_k_per_query=max(1, int(top_k_per_query)), show_per_branch=max(1, int(show_per_branch)), + search_service=search_service, + persist=False, ) report = build_daily_paper_report( search_result=search_result, title=title, top_n=max(1, int(top_n)) diff --git a/src/paperbot/infrastructure/stores/llm_usage_store.py b/src/paperbot/infrastructure/stores/llm_usage_store.py new file mode 100644 index 00000000..ac4892cf --- /dev/null +++ b/src/paperbot/infrastructure/stores/llm_usage_store.py @@ -0,0 +1,157 @@ +from __future__ import annotations + +from collections import defaultdict +from datetime import datetime, timedelta, timezone +from typing import Any, Dict, List, Optional, Tuple + +from sqlalchemy import select + +from paperbot.infrastructure.stores.models import Base, LLMUsageModel +from paperbot.infrastructure.stores.sqlalchemy_db import SessionProvider, get_db_url + + +def _utcnow() -> datetime: + return datetime.now(timezone.utc) + + +class LLMUsageStore: + """Persist and aggregate LLM token/cost usage records.""" + + def __init__(self, db_url: Optional[str] = None, *, auto_create_schema: bool = True): + self.db_url = db_url or get_db_url() + self._provider = SessionProvider(self.db_url) + if auto_create_schema: + Base.metadata.create_all(self._provider.engine) + + def record_usage( + self, + *, + task_type: str, + provider_name: str, + model_name: str, + prompt_tokens: int, + completion_tokens: int, + estimated_cost_usd: float, + metadata: Optional[Dict[str, Any]] = None, + ) -> Dict[str, Any]: + ts = _utcnow() + total_tokens = max(0, int(prompt_tokens) + int(completion_tokens)) + row = LLMUsageModel( + ts=ts, + task_type=(task_type or "default")[:32], + provider_name=(provider_name or "unknown")[:64], + model_name=(model_name or "")[:128], + prompt_tokens=max(0, int(prompt_tokens)), + completion_tokens=max(0, int(completion_tokens)), + total_tokens=total_tokens, + estimated_cost_usd=max(0.0, float(estimated_cost_usd or 0.0)), + metadata_json="{}", + ) + + import json + + row.metadata_json = json.dumps(metadata or {}, ensure_ascii=False) + + with self._provider.session() as session: + session.add(row) + session.commit() + session.refresh(row) + return { + "id": int(row.id), + "ts": row.ts.isoformat() if row.ts else None, + "task_type": row.task_type, + "provider_name": row.provider_name, + "model_name": row.model_name, + "prompt_tokens": int(row.prompt_tokens or 0), + "completion_tokens": int(row.completion_tokens or 0), + "total_tokens": int(row.total_tokens or 0), + "estimated_cost_usd": float(row.estimated_cost_usd or 0.0), + } + + def summarize(self, *, days: int = 7) -> Dict[str, Any]: + window_days = max(1, min(int(days), 90)) + since = _utcnow() - timedelta(days=window_days) + + with self._provider.session() as session: + rows = ( + session.execute(select(LLMUsageModel).where(LLMUsageModel.ts >= since)) + .scalars() + .all() + ) + + daily_map: Dict[str, Dict[str, Any]] = {} + provider_model_map: Dict[Tuple[str, str], Dict[str, Any]] = {} + + for row in rows: + date_key = (row.ts or _utcnow()).date().isoformat() + day = daily_map.setdefault( + date_key, + { + "date": date_key, + "total_tokens": 0, + "total_cost_usd": 0.0, + "providers": defaultdict(int), + }, + ) + provider_key = (row.provider_name or "unknown").strip() or "unknown" + model_key = (row.model_name or "").strip() or "unknown" + total_tokens = int(row.total_tokens or 0) + total_cost = float(row.estimated_cost_usd or 0.0) + + day["total_tokens"] += total_tokens + day["total_cost_usd"] += total_cost + day["providers"][provider_key] += total_tokens + + key = (provider_key, model_key) + bucket = provider_model_map.setdefault( + key, + { + "provider_name": provider_key, + "model_name": model_key, + "calls": 0, + "total_tokens": 0, + "total_cost_usd": 0.0, + }, + ) + bucket["calls"] += 1 + bucket["total_tokens"] += total_tokens + bucket["total_cost_usd"] += total_cost + + daily_rows: List[Dict[str, Any]] = [] + for date_key in sorted(daily_map.keys()): + row = daily_map[date_key] + daily_rows.append( + { + "date": row["date"], + "total_tokens": int(row["total_tokens"]), + "total_cost_usd": round(float(row["total_cost_usd"]), 8), + "providers": dict(row["providers"]), + } + ) + + provider_model_rows = sorted( + provider_model_map.values(), + key=lambda x: int(x["total_tokens"]), + reverse=True, + ) + + totals = { + "calls": int(sum(x["calls"] for x in provider_model_rows)), + "total_tokens": int(sum(x["total_tokens"] for x in provider_model_rows)), + "total_cost_usd": round( + float(sum(x["total_cost_usd"] for x in provider_model_rows)), 8 + ), + } + + return { + "window_days": window_days, + "daily": daily_rows, + "provider_models": provider_model_rows, + "totals": totals, + } + + def close(self) -> None: + try: + self._provider.engine.dispose() + except Exception: + pass diff --git a/src/paperbot/infrastructure/stores/model_endpoint_store.py b/src/paperbot/infrastructure/stores/model_endpoint_store.py new file mode 100644 index 00000000..1d8ac35e --- /dev/null +++ b/src/paperbot/infrastructure/stores/model_endpoint_store.py @@ -0,0 +1,240 @@ +from __future__ import annotations + +import os +from datetime import datetime, timezone +from typing import Any, Dict, List, Optional + +from sqlalchemy import delete, desc, select + +from paperbot.infrastructure.stores.models import Base, ModelEndpointModel +from paperbot.infrastructure.stores.sqlalchemy_db import SessionProvider, get_db_url + +_ALLOWED_VENDORS = { + "openai", + "openai_compatible", + "anthropic", + "ollama", +} + +_ALLOWED_TASK_TYPES = { + "default", + "extraction", + "summary", + "analysis", + "reasoning", + "code", + "review", + "chat", +} + + +def _utcnow() -> datetime: + return datetime.now(timezone.utc) + + +class ModelEndpointStore: + def __init__(self, db_url: Optional[str] = None, *, auto_create_schema: bool = True): + self.db_url = db_url or get_db_url() + self._provider = SessionProvider(self.db_url) + if auto_create_schema: + Base.metadata.create_all(self._provider.engine) + + def list_endpoints( + self, *, enabled_only: bool = False, include_secrets: bool = False + ) -> List[Dict[str, Any]]: + with self._provider.session() as session: + stmt = select(ModelEndpointModel) + if enabled_only: + stmt = stmt.where(ModelEndpointModel.enabled.is_(True)) + stmt = stmt.order_by(desc(ModelEndpointModel.is_default), ModelEndpointModel.name) + rows = session.execute(stmt).scalars().all() + return [self._to_dict(row, include_secrets=include_secrets) for row in rows] + + def get_endpoint( + self, endpoint_id: int, *, include_secrets: bool = False + ) -> Optional[Dict[str, Any]]: + with self._provider.session() as session: + row = session.execute( + select(ModelEndpointModel).where(ModelEndpointModel.id == int(endpoint_id)) + ).scalar_one_or_none() + return self._to_dict(row, include_secrets=include_secrets) if row else None + + def upsert_endpoint( + self, + *, + payload: Dict[str, Any], + endpoint_id: Optional[int] = None, + ) -> Dict[str, Any]: + now = _utcnow() + with self._provider.session() as session: + row: Optional[ModelEndpointModel] = None + if endpoint_id is not None: + row = session.execute( + select(ModelEndpointModel).where(ModelEndpointModel.id == int(endpoint_id)) + ).scalar_one_or_none() + + creating = row is None + if row is None: + row = ModelEndpointModel( + name="", + vendor="openai_compatible", + api_key_env="OPENAI_API_KEY", + models_json="[]", + task_types_json="[]", + enabled=True, + is_default=False, + created_at=now, + updated_at=now, + ) + session.add(row) + + name = str(payload.get("name") or row.name or "").strip() + if not name: + raise ValueError("name is required") + vendor = str(payload.get("vendor") or row.vendor or "openai_compatible").strip().lower() + if vendor not in _ALLOWED_VENDORS: + raise ValueError(f"unsupported vendor: {vendor}") + + models = payload.get("models") + if models is None: + models = row.get_models() + if not isinstance(models, list): + models = [str(models)] + normalized_models = [str(x).strip() for x in models if str(x).strip()] + if not normalized_models: + raise ValueError("at least one model is required") + + task_types = payload.get("task_types") + if task_types is None: + task_types = row.get_task_types() + if not isinstance(task_types, list): + task_types = [str(task_types)] + normalized_tasks = sorted( + { + str(x).strip().lower() + for x in task_types + if str(x).strip().lower() in _ALLOWED_TASK_TYPES + } + ) + + row.name = name + row.vendor = vendor + row.base_url = str(payload.get("base_url") or row.base_url or "").strip() or None + row.api_key_env = ( + str(payload.get("api_key_env") or row.api_key_env or "OPENAI_API_KEY").strip() + or "OPENAI_API_KEY" + ) + if "api_key" in payload: + api_key_text = str(payload.get("api_key") or "").strip() + if not api_key_text: + row.api_key_value = None + elif not api_key_text.startswith("***"): + row.api_key_value = api_key_text + row.enabled = bool(payload.get("enabled", row.enabled)) + row.is_default = bool(payload.get("is_default", row.is_default)) + row.set_models(normalized_models) + row.set_task_types(normalized_tasks) + row.updated_at = now + if creating: + row.created_at = now + + session.flush() + + if row.is_default: + session.execute( + ModelEndpointModel.__table__.update() + .where(ModelEndpointModel.id != row.id) + .values(is_default=False, updated_at=now) + ) + elif not session.execute( + select(ModelEndpointModel).where(ModelEndpointModel.is_default.is_(True)).limit(1) + ).scalar_one_or_none(): + # Ensure there is always one default endpoint for fallback routing. + row.is_default = True + + session.commit() + session.refresh(row) + return self._to_dict(row) + + def delete_endpoint(self, endpoint_id: int) -> bool: + with self._provider.session() as session: + row = session.execute( + select(ModelEndpointModel).where(ModelEndpointModel.id == int(endpoint_id)) + ).scalar_one_or_none() + if row is None: + return False + deleted_default = bool(row.is_default) + session.execute( + delete(ModelEndpointModel).where(ModelEndpointModel.id == int(endpoint_id)) + ) + session.flush() + + if deleted_default: + replacement = session.execute( + select(ModelEndpointModel).order_by(ModelEndpointModel.id.asc()).limit(1) + ).scalar_one_or_none() + if replacement is not None: + replacement.is_default = True + replacement.updated_at = _utcnow() + session.add(replacement) + + session.commit() + return True + + def activate_endpoint(self, endpoint_id: int) -> Optional[Dict[str, Any]]: + now = _utcnow() + with self._provider.session() as session: + row = session.execute( + select(ModelEndpointModel).where(ModelEndpointModel.id == int(endpoint_id)) + ).scalar_one_or_none() + if row is None: + return None + + session.execute( + ModelEndpointModel.__table__.update().values(is_default=False, updated_at=now) + ) + row.is_default = True + row.enabled = True + row.updated_at = now + session.add(row) + session.commit() + session.refresh(row) + return self._to_dict(row) + + @staticmethod + def _to_dict(row: ModelEndpointModel, *, include_secrets: bool = False) -> Dict[str, Any]: + models = row.get_models() + task_types = row.get_task_types() + key_raw = str(row.api_key_value or "").strip() + key_present = bool(key_raw) or bool(os.getenv(row.api_key_env or "")) + key_display = key_raw if include_secrets else _mask_secret(key_raw) + return { + "id": int(row.id), + "name": row.name, + "vendor": row.vendor, + "base_url": row.base_url, + "api_key_env": row.api_key_env, + "api_key": key_display, + "models": models, + "task_types": task_types, + "enabled": bool(row.enabled), + "is_default": bool(row.is_default), + "api_key_present": key_present, + "created_at": row.created_at.isoformat() if row.created_at else None, + "updated_at": row.updated_at.isoformat() if row.updated_at else None, + } + + def close(self) -> None: + try: + self._provider.engine.dispose() + except Exception: + pass + + +def _mask_secret(value: str) -> str: + text = str(value or "") + if not text: + return "" + if len(text) <= 8: + return "***" + return f"***{text[-8:]}" diff --git a/src/paperbot/infrastructure/stores/models.py b/src/paperbot/infrastructure/stores/models.py index 68819b02..dbdb4cc8 100644 --- a/src/paperbot/infrastructure/stores/models.py +++ b/src/paperbot/infrastructure/stores/models.py @@ -467,9 +467,7 @@ class PaperFeedbackModel(Base): action: Mapped[str] = mapped_column(String(16), index=True) # like/dislike/skip/save/cite # Canonical FK (dual-write migration — will replace paper_id + paper_ref_id) - canonical_paper_id: Mapped[Optional[int]] = mapped_column( - Integer, nullable=True, index=True - ) + canonical_paper_id: Mapped[Optional[int]] = mapped_column(Integer, nullable=True, index=True) weight: Mapped[float] = mapped_column(Float, default=0.0) ts: Mapped[datetime] = mapped_column(DateTime(timezone=True), index=True) @@ -749,7 +747,9 @@ class PaperModel(Base): judge_scores = relationship("PaperJudgeScoreModel", back_populates="paper") reading_status_rows = relationship("PaperReadingStatusModel", back_populates="paper") repo_rows = relationship("PaperRepoModel", back_populates="paper") - identifiers = relationship("PaperIdentifierModel", back_populates="paper", cascade="all, delete-orphan") + identifiers = relationship( + "PaperIdentifierModel", back_populates="paper", cascade="all, delete-orphan" + ) def get_authors(self) -> list: try: @@ -804,6 +804,89 @@ class PaperIdentifierModel(Base): paper = relationship("PaperModel", back_populates="identifiers") +class ModelEndpointModel(Base): + """User-managed LLM provider endpoints for gateway routing.""" + + __tablename__ = "model_endpoints" + __table_args__ = (UniqueConstraint("name", name="uq_model_endpoints_name"),) + + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + name: Mapped[str] = mapped_column(String(64), index=True) + vendor: Mapped[str] = mapped_column(String(32), default="openai_compatible", index=True) + base_url: Mapped[Optional[str]] = mapped_column(String(512), nullable=True) + api_key_env: Mapped[str] = mapped_column(String(64), default="OPENAI_API_KEY") + api_key_value: Mapped[Optional[str]] = mapped_column(String(512), nullable=True) + models_json: Mapped[str] = mapped_column(Text, default="[]") + task_types_json: Mapped[str] = mapped_column(Text, default="[]") + enabled: Mapped[bool] = mapped_column(Boolean, default=True, index=True) + is_default: Mapped[bool] = mapped_column(Boolean, default=False, index=True) + created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), index=True) + updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), index=True) + + def get_models(self) -> list[str]: + try: + rows = json.loads(self.models_json or "[]") + if isinstance(rows, list): + return [str(x).strip() for x in rows if str(x).strip()] + except Exception: + pass + return [] + + def set_models(self, rows: Optional[list[str]]) -> None: + self.models_json = json.dumps( + [str(x).strip() for x in (rows or []) if str(x).strip()], + ensure_ascii=False, + ) + + def get_task_types(self) -> list[str]: + try: + rows = json.loads(self.task_types_json or "[]") + if isinstance(rows, list): + return [str(x).strip() for x in rows if str(x).strip()] + except Exception: + pass + return [] + + def set_task_types(self, rows: Optional[list[str]]) -> None: + self.task_types_json = json.dumps( + [str(x).strip() for x in (rows or []) if str(x).strip()], + ensure_ascii=False, + ) + + +class LLMUsageModel(Base): + """LLM token/cost usage records for dashboard and alerting.""" + + __tablename__ = "llm_usage" + + id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True) + ts: Mapped[datetime] = mapped_column(DateTime(timezone=True), index=True) + task_type: Mapped[str] = mapped_column(String(32), default="default", index=True) + provider_name: Mapped[str] = mapped_column(String(64), default="unknown", index=True) + model_name: Mapped[str] = mapped_column(String(128), default="", index=True) + prompt_tokens: Mapped[int] = mapped_column(Integer, default=0) + completion_tokens: Mapped[int] = mapped_column(Integer, default=0) + total_tokens: Mapped[int] = mapped_column(Integer, default=0, index=True) + estimated_cost_usd: Mapped[float] = mapped_column(Float, default=0.0) + metadata_json: Mapped[str] = mapped_column(Text, default="{}") + + +class PipelineSessionModel(Base): + """Long-running pipeline session checkpoints for resume/recovery.""" + + __tablename__ = "pipeline_sessions" + + session_id: Mapped[str] = mapped_column(String(64), primary_key=True) + workflow: Mapped[str] = mapped_column(String(64), default="", index=True) + status: Mapped[str] = mapped_column(String(32), default="running", index=True) + checkpoint: Mapped[str] = mapped_column(String(64), default="init", index=True) + payload_json: Mapped[str] = mapped_column(Text, default="{}") + state_json: Mapped[str] = mapped_column(Text, default="{}") + result_json: Mapped[str] = mapped_column(Text, default="{}") + created_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), index=True) + updated_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), index=True) + + class HarvestRunModel(Base): """Harvest execution tracking.""" diff --git a/src/paperbot/infrastructure/stores/paper_store.py b/src/paperbot/infrastructure/stores/paper_store.py index 7db8d1ed..0d056062 100644 --- a/src/paperbot/infrastructure/stores/paper_store.py +++ b/src/paperbot/infrastructure/stores/paper_store.py @@ -719,10 +719,14 @@ def get_user_library( if USE_CANONICAL_FK: return self._get_user_library_canonical( - user_id, session_provider=self._provider, - track_id=track_id, actions=actions, - sort_by=sort_by, sort_order=sort_order, - limit=limit, offset=offset, + user_id, + session_provider=self._provider, + track_id=track_id, + actions=actions, + sort_by=sort_by, + sort_order=sort_order, + limit=limit, + offset=offset, ) with self._provider.session() as session: @@ -739,7 +743,12 @@ def get_user_library( # This avoids CAST errors on PostgreSQL for non-numeric paper_ids # Also check library_paper_id from metadata if available base_stmt = ( - select(PaperModel, PaperFeedbackModel) + select( + PaperModel, + PaperFeedbackModel.ts, + PaperFeedbackModel.track_id, + PaperFeedbackModel.action, + ) .join( PaperFeedbackModel, or_( @@ -769,12 +778,20 @@ def get_user_library( # Deduplicate by paper.id, keeping the one with latest timestamp Logger.info("Deduplicating results by paper id", file=LogFiles.HARVEST) - paper_map: Dict[int, Tuple[PaperModel, PaperFeedbackModel]] = {} + paper_map: Dict[int, Tuple[PaperModel, Optional[datetime], Optional[int], str]] = {} for row in all_results: paper = row[0] - feedback = row[1] - if paper.id not in paper_map or feedback.ts > paper_map[paper.id][1].ts: - paper_map[paper.id] = (paper, feedback) + ts = row[1] + fb_track_id = row[2] + fb_action = row[3] + current_ts = ts or datetime.min.replace(tzinfo=timezone.utc) + existing_ts = ( + paper_map[paper.id][1] or datetime.min.replace(tzinfo=timezone.utc) + if paper.id in paper_map + else datetime.min.replace(tzinfo=timezone.utc) + ) + if paper.id not in paper_map or current_ts > existing_ts: + paper_map[paper.id] = (paper, ts, fb_track_id, fb_action) # Convert to list and sort unique_results = list(paper_map.values()) @@ -814,9 +831,9 @@ def get_user_library( return [ LibraryPaper( paper=row[0], - saved_at=row[1].ts, - track_id=row[1].track_id, - action=row[1].action, + saved_at=row[1], + track_id=row[2], + action=row[3], ) for row in paginated_results ], total @@ -895,6 +912,37 @@ def remove_from_library(self, user_id: str, paper_id: int) -> bool: session.commit() return result.rowcount > 0 + def get_latest_judge_scores(self, paper_ids: List[int]) -> Dict[int, Dict[str, Any]]: + """Fetch latest judge score per paper id.""" + ids = sorted({int(pid) for pid in paper_ids if int(pid) > 0}) + if not ids: + return {} + + with self._provider.session() as session: + rows = ( + session.execute( + select(PaperJudgeScoreModel) + .where(PaperJudgeScoreModel.paper_id.in_(ids)) + .order_by(desc(PaperJudgeScoreModel.scored_at), desc(PaperJudgeScoreModel.id)) + ) + .scalars() + .all() + ) + + latest: Dict[int, Dict[str, Any]] = {} + for row in rows: + pid = int(row.paper_id) + if pid in latest: + continue + latest[pid] = { + "overall": float(row.overall or 0.0), + "recommendation": str(row.recommendation or ""), + "one_line_summary": str(row.one_line_summary or ""), + "judge_model": str(row.judge_model or ""), + "scored_at": row.scored_at.isoformat() if row.scored_at else None, + } + return latest + def create_harvest_run( self, run_id: str, diff --git a/src/paperbot/infrastructure/stores/pipeline_session_store.py b/src/paperbot/infrastructure/stores/pipeline_session_store.py new file mode 100644 index 00000000..a972e795 --- /dev/null +++ b/src/paperbot/infrastructure/stores/pipeline_session_store.py @@ -0,0 +1,243 @@ +from __future__ import annotations + +import json +from datetime import datetime, timezone +from typing import Any, Dict, Optional +from uuid import uuid4 + +from sqlalchemy import desc, select + +from paperbot.infrastructure.stores.models import Base, PipelineSessionModel +from paperbot.infrastructure.stores.sqlalchemy_db import SessionProvider, get_db_url + + +def _utcnow() -> datetime: + return datetime.now(timezone.utc) + + +class PipelineSessionStore: + """Persist lightweight workflow checkpoints for resume support.""" + + def __init__(self, db_url: Optional[str] = None, *, auto_create_schema: bool = True): + self.db_url = db_url or get_db_url() + self._provider = SessionProvider(self.db_url) + if auto_create_schema: + Base.metadata.create_all(self._provider.engine) + + def start_session( + self, + *, + workflow: str, + payload: Optional[Dict[str, Any]] = None, + session_id: Optional[str] = None, + resume: bool = False, + ) -> Dict[str, Any]: + resolved_id = (session_id or "").strip() or uuid4().hex + now = _utcnow() + + with self._provider.session() as session: + row = session.execute( + select(PipelineSessionModel).where(PipelineSessionModel.session_id == resolved_id) + ).scalar_one_or_none() + + if row and resume: + row.updated_at = now + session.add(row) + session.commit() + return self._to_dict(row) + + if row is None: + row = PipelineSessionModel( + session_id=resolved_id, + workflow=(workflow or "")[:64], + status="running", + checkpoint="init", + payload_json=json.dumps(payload or {}, ensure_ascii=False), + state_json="{}", + result_json="{}", + created_at=now, + updated_at=now, + ) + session.add(row) + else: + row.workflow = (workflow or row.workflow or "")[:64] + row.status = "running" + row.checkpoint = "init" + row.payload_json = json.dumps(payload or {}, ensure_ascii=False) + row.state_json = "{}" + row.result_json = "{}" + row.updated_at = now + session.add(row) + + session.commit() + session.refresh(row) + return self._to_dict(row) + + def get_session(self, session_id: str) -> Optional[Dict[str, Any]]: + sid = (session_id or "").strip() + if not sid: + return None + with self._provider.session() as session: + row = session.execute( + select(PipelineSessionModel).where(PipelineSessionModel.session_id == sid) + ).scalar_one_or_none() + return self._to_dict(row) if row else None + + def list_sessions( + self, + *, + workflow: Optional[str] = None, + status: Optional[str] = None, + limit: int = 20, + ) -> list[Dict[str, Any]]: + with self._provider.session() as session: + stmt = select(PipelineSessionModel) + if workflow: + stmt = stmt.where(PipelineSessionModel.workflow == str(workflow)[:64]) + if status: + stmt = stmt.where(PipelineSessionModel.status == str(status)[:32]) + stmt = stmt.order_by(desc(PipelineSessionModel.updated_at)).limit(max(1, int(limit))) + rows = session.execute(stmt).scalars().all() + return [self._to_dict(row) for row in rows] + + def save_checkpoint( + self, + *, + session_id: str, + checkpoint: str, + state: Optional[Dict[str, Any]] = None, + ) -> Optional[Dict[str, Any]]: + sid = (session_id or "").strip() + if not sid: + return None + + with self._provider.session() as session: + row = session.execute( + select(PipelineSessionModel).where(PipelineSessionModel.session_id == sid) + ).scalar_one_or_none() + if row is None: + return None + + row.checkpoint = (checkpoint or "")[:64] or row.checkpoint + row.state_json = json.dumps(state or {}, ensure_ascii=False) + row.status = "running" + row.updated_at = _utcnow() + session.add(row) + session.commit() + session.refresh(row) + return self._to_dict(row) + + def save_result( + self, + *, + session_id: str, + result: Optional[Dict[str, Any]] = None, + status: str = "completed", + ) -> Optional[Dict[str, Any]]: + sid = (session_id or "").strip() + if not sid: + return None + + with self._provider.session() as session: + row = session.execute( + select(PipelineSessionModel).where(PipelineSessionModel.session_id == sid) + ).scalar_one_or_none() + if row is None: + return None + + row.status = (status or "completed")[:32] + row.checkpoint = "result" + row.result_json = json.dumps(result or {}, ensure_ascii=False) + row.updated_at = _utcnow() + session.add(row) + session.commit() + session.refresh(row) + return self._to_dict(row) + + def mark_failed(self, *, session_id: str, error: str) -> Optional[Dict[str, Any]]: + sid = (session_id or "").strip() + if not sid: + return None + + with self._provider.session() as session: + row = session.execute( + select(PipelineSessionModel).where(PipelineSessionModel.session_id == sid) + ).scalar_one_or_none() + if row is None: + return None + + state = _safe_json_dict(row.state_json) + state["error"] = str(error or "") + row.status = "failed" + row.state_json = json.dumps(state, ensure_ascii=False) + row.updated_at = _utcnow() + session.add(row) + session.commit() + session.refresh(row) + return self._to_dict(row) + + def update_status( + self, + *, + session_id: str, + status: str, + checkpoint: Optional[str] = None, + state_patch: Optional[Dict[str, Any]] = None, + result: Optional[Dict[str, Any]] = None, + ) -> Optional[Dict[str, Any]]: + sid = (session_id or "").strip() + if not sid: + return None + + with self._provider.session() as session: + row = session.execute( + select(PipelineSessionModel).where(PipelineSessionModel.session_id == sid) + ).scalar_one_or_none() + if row is None: + return None + + row.status = (status or row.status or "running")[:32] + if checkpoint: + row.checkpoint = str(checkpoint)[:64] + + if state_patch: + merged = _safe_json_dict(row.state_json) + merged.update(state_patch) + row.state_json = json.dumps(merged, ensure_ascii=False) + + if result is not None: + row.result_json = json.dumps(result, ensure_ascii=False) + + row.updated_at = _utcnow() + session.add(row) + session.commit() + session.refresh(row) + return self._to_dict(row) + + @staticmethod + def _to_dict(row: PipelineSessionModel) -> Dict[str, Any]: + return { + "session_id": row.session_id, + "workflow": row.workflow, + "status": row.status, + "checkpoint": row.checkpoint, + "payload": _safe_json_dict(row.payload_json), + "state": _safe_json_dict(row.state_json), + "result": _safe_json_dict(row.result_json), + "created_at": row.created_at.isoformat() if row.created_at else None, + "updated_at": row.updated_at.isoformat() if row.updated_at else None, + } + + def close(self) -> None: + try: + self._provider.engine.dispose() + except Exception: + pass + + +def _safe_json_dict(raw: str) -> Dict[str, Any]: + try: + data = json.loads(raw or "{}") + return data if isinstance(data, dict) else {} + except Exception: + return {} diff --git a/src/paperbot/infrastructure/stores/research_store.py b/src/paperbot/infrastructure/stores/research_store.py index 597a5644..ad9d3e72 100644 --- a/src/paperbot/infrastructure/stores/research_store.py +++ b/src/paperbot/infrastructure/stores/research_store.py @@ -8,6 +8,7 @@ from sqlalchemy import desc, func, or_, select from sqlalchemy.exc import IntegrityError +from paperbot.application.services.identity_resolver import IdentityResolver from paperbot.domain.paper_identity import normalize_arxiv_id, normalize_doi from paperbot.infrastructure.stores.models import ( Base, @@ -91,6 +92,7 @@ class SqlAlchemyResearchStore: def __init__(self, db_url: Optional[str] = None, *, auto_create_schema: bool = True): self.db_url = db_url or get_db_url() self._provider = SessionProvider(self.db_url) + self._identity_resolver = IdentityResolver(db_url=self.db_url) if auto_create_schema: Base.metadata.create_all(self._provider.engine) @@ -529,6 +531,7 @@ def list_saved_papers( self, *, user_id: str, + track_id: Optional[int] = None, limit: int = 200, sort_by: str = "saved_at", ) -> List[Dict[str, Any]]: @@ -555,6 +558,11 @@ def list_saved_papers( PaperFeedbackModel.user_id == user_id, PaperFeedbackModel.action == "save", PaperFeedbackModel.paper_ref_id.is_not(None), + ( + PaperFeedbackModel.track_id == int(track_id) + if track_id is not None + else True + ), ) ) .scalars() @@ -635,6 +643,212 @@ def list_saved_papers( return rows[: max(1, int(limit))] + def list_track_feed( + self, + *, + user_id: str, + track_id: int, + limit: int = 20, + offset: int = 0, + ) -> Dict[str, Any]: + with self._provider.session() as session: + track = session.execute( + select(ResearchTrackModel).where( + ResearchTrackModel.user_id == user_id, + ResearchTrackModel.id == int(track_id), + ResearchTrackModel.archived_at.is_(None), + ) + ).scalar_one_or_none() + if track is None: + return {"items": [], "total": 0} + + track_dict = self._track_to_dict(track) + raw_terms = [ + *track_dict.get("keywords", []), + *track_dict.get("methods", []), + *track_dict.get("venues", []), + ] + terms = sorted({str(term).strip().lower() for term in raw_terms if str(term).strip()}) + + stmt = select(PaperModel).where(PaperModel.deleted_at.is_(None)) + if terms: + term_filters = [] + for term in terms: + like = f"%{term}%" + term_filters.extend( + [ + func.lower(func.coalesce(PaperModel.title, "")).like(like), + func.lower(func.coalesce(PaperModel.abstract, "")).like(like), + func.lower(func.coalesce(PaperModel.venue, "")).like(like), + func.lower(func.coalesce(PaperModel.keywords_json, "")).like(like), + func.lower(func.coalesce(PaperModel.fields_of_study_json, "")).like( + like + ), + ] + ) + stmt = stmt.where(or_(*term_filters)) + + feedback_rows = ( + session.execute( + select(PaperFeedbackModel) + .where( + PaperFeedbackModel.user_id == user_id, + PaperFeedbackModel.track_id == int(track_id), + PaperFeedbackModel.action.in_(["save", "like", "dislike", "skip"]), + ) + .order_by(desc(PaperFeedbackModel.ts), desc(PaperFeedbackModel.id)) + ) + .scalars() + .all() + ) + + feedback_candidate_ids = { + int(row.canonical_paper_id or row.paper_ref_id or 0) + for row in feedback_rows + if int(row.canonical_paper_id or row.paper_ref_id or 0) > 0 + } + + fetch_cap = max(200, (int(offset) + int(limit)) * 8) + candidates = ( + session.execute( + stmt.order_by(desc(PaperModel.created_at), desc(PaperModel.id)).limit(fetch_cap) + ) + .scalars() + .all() + ) + + if feedback_candidate_ids: + existing_ids = {int(p.id) for p in candidates} + missing_ids = sorted(feedback_candidate_ids - existing_ids) + if missing_ids: + extra = ( + session.execute( + select(PaperModel) + .where(PaperModel.id.in_(missing_ids), PaperModel.deleted_at.is_(None)) + .order_by(desc(PaperModel.created_at), desc(PaperModel.id)) + ) + .scalars() + .all() + ) + candidates.extend(extra) + + if not candidates: + return {"items": [], "total": 0} + + candidate_ids = [int(p.id) for p in candidates] + + feedback_by_paper: Dict[int, PaperFeedbackModel] = {} + feedback_summary_by_paper: Dict[int, Dict[str, int]] = {} + for row in feedback_rows: + pid = int(row.canonical_paper_id or row.paper_ref_id or 0) + if pid <= 0: + continue + action = str(row.action or "").strip().lower() + if action: + action_counter = feedback_summary_by_paper.setdefault(pid, {}) + action_counter[action] = action_counter.get(action, 0) + 1 + if pid not in feedback_by_paper: + feedback_by_paper[pid] = row + + status_by_paper = { + int(row.paper_id): row + for row in session.execute( + select(PaperReadingStatusModel).where( + PaperReadingStatusModel.user_id == user_id, + PaperReadingStatusModel.paper_id.in_(candidate_ids), + ) + ) + .scalars() + .all() + } + + latest_judge_by_paper: Dict[int, PaperJudgeScoreModel] = {} + judge_rows = ( + session.execute( + select(PaperJudgeScoreModel) + .where(PaperJudgeScoreModel.paper_id.in_(candidate_ids)) + .order_by(desc(PaperJudgeScoreModel.scored_at), desc(PaperJudgeScoreModel.id)) + ) + .scalars() + .all() + ) + for judge in judge_rows: + pid = int(judge.paper_id or 0) + if pid > 0 and pid not in latest_judge_by_paper: + latest_judge_by_paper[pid] = judge + + scored_rows: List[Dict[str, Any]] = [] + for paper in candidates: + pid = int(paper.id) + text_blob = " ".join( + [ + str(paper.title or ""), + str(paper.abstract or ""), + str(paper.venue or ""), + " ".join(str(x) for x in (paper.get_keywords() or [])), + " ".join(str(x) for x in (paper.get_fields_of_study() or [])), + ] + ).lower() + + matched_terms = [term for term in terms if term and term in text_blob] + keyword_score = float(len(matched_terms)) + + latest_feedback = feedback_by_paper.get(pid) + latest_feedback_action = ( + str(latest_feedback.action or "").strip().lower() if latest_feedback else "" + ) + feedback_boost = { + "save": 3.0, + "like": 2.0, + "skip": -1.0, + "dislike": -4.0, + }.get(latest_feedback_action, 0.0) + + citation_score = min(float(paper.citation_count or 0) / 200.0, 2.0) + judge_row = latest_judge_by_paper.get(pid) + judge_score = float(judge_row.overall or 0.0) if judge_row else 0.0 + + if terms and keyword_score <= 0 and abs(feedback_boost) < 1e-6: + continue + + feed_score = ( + keyword_score * 2.5 + feedback_boost + citation_score + judge_score * 0.3 + ) + scored_rows.append( + { + "paper": self._paper_to_dict(paper), + "latest_judge": ( + self._judge_score_to_dict(latest_judge_by_paper[pid]) + if pid in latest_judge_by_paper + else None + ), + "reading_status": ( + self._reading_status_to_dict(status_by_paper[pid]) + if pid in status_by_paper + else None + ), + "latest_feedback_action": latest_feedback_action or None, + "feedback_summary": feedback_summary_by_paper.get(pid, {}), + "matched_terms": matched_terms, + "keyword_score": keyword_score, + "feed_score": round(feed_score, 4), + } + ) + + scored_rows.sort( + key=lambda row: ( + float(row.get("feed_score") or 0.0), + float(((row.get("latest_judge") or {}).get("overall") or 0.0)), + str(((row.get("paper") or {}).get("created_at") or "")), + ), + reverse=True, + ) + + total = len(scored_rows) + start = max(0, int(offset)) + end = start + max(1, int(limit)) + return {"items": scored_rows[start:end], "total": total} + def ingest_repo_enrichment_rows( self, *, @@ -1321,8 +1535,29 @@ def _upsert_paper_repo_row( session.add(row) return created - @staticmethod def _resolve_paper_ref_id( + self, + *, + session, + paper_id: str, + metadata: Dict[str, Any], + ) -> Optional[int]: + pid = (paper_id or "").strip() + hints = dict(metadata or {}) + + # Main path: centralized identity resolver (paper_identifiers + normalized fallbacks). + resolved = self._identity_resolver.resolve(pid, hints=hints) + if resolved is not None: + return int(resolved) + + Logger.info( + "IdentityResolver miss; falling back to legacy paper_id resolution", + file=LogFiles.HARVEST, + ) + return self._resolve_paper_ref_id_legacy(session=session, paper_id=pid, metadata=hints) + + @staticmethod + def _resolve_paper_ref_id_legacy( *, session, paper_id: str, diff --git a/src/paperbot/presentation/cli/main.py b/src/paperbot/presentation/cli/main.py index 39d9019a..e3e138b1 100644 --- a/src/paperbot/presentation/cli/main.py +++ b/src/paperbot/presentation/cli/main.py @@ -26,6 +26,12 @@ render_daily_paper_markdown, ) from paperbot.application.services.daily_push_service import DailyPushService +from paperbot.application.workflows.unified_topic_search import ( + make_default_search_service, + run_unified_topic_search, +) +from paperbot.infrastructure.stores.paper_store import PaperStore +from paperbot.infrastructure.stores.pipeline_session_store import PipelineSessionStore # Load local .env automatically for CLI workflows using LLM providers. load_dotenv(find_dotenv(usecwd=True), override=False) @@ -158,6 +164,16 @@ def create_parser() -> argparse.ArgumentParser: default=None, help="推送渠道:email/slack/dingding(可重复指定)", ) + daily_parser.add_argument( + "--session-id", + default=None, + help="会话 ID(用于断点恢复;不传则自动生成)", + ) + daily_parser.add_argument( + "--resume", + action="store_true", + help="从最近 checkpoint 恢复执行", + ) # version parser.add_argument("--version", "-v", action="store_true", help="显示版本") @@ -241,10 +257,14 @@ async def _quick_score(paper_id: str): print("Please use the main.py entry point instead.") -def _create_topic_search_workflow(): - from paperbot.application.workflows.paperscool_topic_search import PapersCoolTopicSearchWorkflow +_paper_search_service = None + - return PapersCoolTopicSearchWorkflow() +def _get_paper_search_service(): + global _paper_search_service + if _paper_search_service is None: + _paper_search_service = make_default_search_service(registry=PaperStore()) + return _paper_search_service def _run_topic_search(parsed: argparse.Namespace) -> int: @@ -255,13 +275,16 @@ def _run_topic_search(parsed: argparse.Namespace) -> int: branches = parsed.branches or ["arxiv", "venue"] sources = parsed.sources or ["papers_cool"] - workflow = _create_topic_search_workflow() - result = workflow.run( - queries=queries, - sources=sources, - branches=branches, - top_k_per_query=max(1, int(parsed.top_k)), - show_per_branch=max(1, int(parsed.show)), + result = asyncio.run( + run_unified_topic_search( + queries=queries, + sources=sources, + branches=branches, + top_k_per_query=max(1, int(parsed.top_k)), + show_per_branch=max(1, int(parsed.show)), + search_service=_get_paper_search_service(), + persist=False, + ) ) if parsed.json: @@ -287,21 +310,55 @@ def _run_daily_paper(parsed: argparse.Namespace) -> int: branches = parsed.branches or ["arxiv", "venue"] sources = parsed.sources or ["papers_cool"] - workflow = _create_topic_search_workflow() - effective_top_k = max(1, int(parsed.top_k), int(parsed.top_n)) - search_result = workflow.run( - queries=queries, - sources=sources, - branches=branches, - top_k_per_query=effective_top_k, - show_per_branch=max(1, int(parsed.show)), + session_store = PipelineSessionStore() + session = session_store.start_session( + workflow="cli_daily_paper", + payload={k: getattr(parsed, k) for k in vars(parsed)}, + session_id=getattr(parsed, "session_id", None), + resume=bool(getattr(parsed, "resume", False)), ) + session_id = str(session.get("session_id") or "") + state: dict = session.get("state") if getattr(parsed, "resume", False) else {} + + if getattr(parsed, "resume", False) and session.get("status") == "completed": + done = session.get("result") if isinstance(session.get("result"), dict) else {} + if done: + if parsed.json: + print(json.dumps(done, ensure_ascii=False, indent=2)) + else: + print(f"resumed session: {session_id}") + print("session already completed; returning cached result") + return 0 + + if isinstance(state.get("report"), dict): + report = dict(state.get("report") or {}) + search_result = ( + state.get("search_result") if isinstance(state.get("search_result"), dict) else {} + ) + else: + effective_top_k = max(1, int(parsed.top_k), int(parsed.top_n)) + search_result = asyncio.run( + run_unified_topic_search( + queries=queries, + sources=sources, + branches=branches, + top_k_per_query=effective_top_k, + show_per_branch=max(1, int(parsed.show)), + search_service=_get_paper_search_service(), + persist=False, + ) + ) - report = build_daily_paper_report( - search_result=search_result, - title=parsed.title, - top_n=max(1, int(parsed.top_n)), - ) + report = build_daily_paper_report( + search_result=search_result, + title=parsed.title, + top_n=max(1, int(parsed.top_n)), + ) + session_store.save_checkpoint( + session_id=session_id, + checkpoint="report_built", + state={"search_result": search_result, "report": report}, + ) llm_enabled = bool(parsed.with_llm) llm_features = normalize_llm_features(parsed.llm_features or ["summary"]) if llm_enabled: @@ -356,8 +413,22 @@ def _run_daily_paper(parsed: argparse.Namespace) -> int: channels_override=parsed.notify_channels, ) + session_store.save_result( + session_id=session_id, + status="completed", + result={ + "session_id": session_id, + "report": report, + "markdown": markdown, + "markdown_path": markdown_path, + "json_path": json_path, + "notify": notify_result, + }, + ) + if parsed.json: payload = { + "session_id": session_id, "report": report, "markdown": markdown, "markdown_path": markdown_path, diff --git a/tests/unit/test_context_engine_enrichment.py b/tests/unit/test_context_engine_enrichment.py new file mode 100644 index 00000000..21480030 --- /dev/null +++ b/tests/unit/test_context_engine_enrichment.py @@ -0,0 +1,63 @@ +from __future__ import annotations + +from paperbot.context_engine.engine import ContextEngine + + +class _FakePaperStore: + def __init__(self): + self.calls = [] + + def get_latest_judge_scores(self, paper_ids): + self.calls.append(list(paper_ids)) + return { + 1: { + "overall": 4.6, + "recommendation": "must_read", + "one_line_summary": "Strong relevance", + "judge_model": "judge-x", + "scored_at": "2026-02-12T00:00:00+00:00", + } + } + + +def test_attach_latest_judge_only_for_numeric_ids(): + store = _FakePaperStore() + engine = ContextEngine( + research_store=object(), + memory_store=object(), + paper_store=store, + track_router=object(), + ) + papers = [ + {"paper_id": "1", "title": "A"}, + {"paper_id": "not-numeric", "title": "B"}, + {"paper_id": "", "title": "C"}, + ] + + engine._attach_latest_judge(papers) + + assert store.calls == [[1]] + assert papers[0]["latest_judge"]["overall"] == 4.6 + assert "latest_judge" not in papers[1] + assert "latest_judge" not in papers[2] + + +def test_attach_feedback_flags_marks_saved_and_liked(): + papers = [ + {"paper_id": "11", "title": "Saved only"}, + {"paper_id": "22", "title": "Saved + liked"}, + {"paper_id": "33", "title": "None"}, + ] + + ContextEngine._attach_feedback_flags( + papers, + saved_ids={"11", "22"}, + liked_ids={"22"}, + ) + + assert papers[0].get("is_saved") is True + assert papers[0].get("is_liked") is None + assert papers[1].get("is_saved") is True + assert papers[1].get("is_liked") is True + assert papers[2].get("is_saved") is None + assert papers[2].get("is_liked") is None diff --git a/tests/unit/test_llm_service.py b/tests/unit/test_llm_service.py index 9856b9c3..1bce26ef 100644 --- a/tests/unit/test_llm_service.py +++ b/tests/unit/test_llm_service.py @@ -36,6 +36,25 @@ def get_provider(self, task_type: str = "default"): return self.provider +class _FakeResolver: + def __init__(self, provider: _FakeProvider): + self.provider = provider + self.task_types = [] + + def get_provider(self, task_type: str = "default"): + self.task_types.append(task_type) + return self.provider + + +class _FakeUsageStore: + def __init__(self): + self.rows = [] + + def record_usage(self, **kwargs): + self.rows.append(kwargs) + return {"id": len(self.rows)} + + def test_complete_uses_cache_for_same_request(): provider = _FakeProvider(response="cached") service = LLMService(router=_FakeRouter(provider)) @@ -85,3 +104,31 @@ def test_describe_task_provider_returns_model_metadata(): assert info["provider_name"] == "fake" assert info["model_name"] == "fake-model" + + +def test_provider_resolver_can_be_injected(): + provider = _FakeProvider(response="resolver-ok") + resolver = _FakeResolver(provider) + service = LLMService(provider_resolver=resolver) + + output = service.complete(task_type="chat", system="sys", user="hello") + + assert output == "resolver-ok" + assert resolver.task_types == ["chat"] + + +def test_complete_records_usage(): + provider = _FakeProvider(response="usage") + usage_store = _FakeUsageStore() + service = LLMService(router=_FakeRouter(provider), usage_store=usage_store) + + result = service.complete(task_type="summary", system="system prompt", user="user prompt") + + assert result == "usage" + assert len(usage_store.rows) == 1 + row = usage_store.rows[0] + assert row["task_type"] == "summary" + assert row["provider_name"] == "fake" + assert row["prompt_tokens"] >= 1 + assert row["completion_tokens"] >= 1 + assert row["estimated_cost_usd"] >= 0.0 diff --git a/tests/unit/test_llm_usage_store.py b/tests/unit/test_llm_usage_store.py new file mode 100644 index 00000000..ecdc46cf --- /dev/null +++ b/tests/unit/test_llm_usage_store.py @@ -0,0 +1,38 @@ +from __future__ import annotations + +from pathlib import Path + +from paperbot.infrastructure.stores.llm_usage_store import LLMUsageStore + + +def test_llm_usage_store_records_and_summarizes(tmp_path: Path): + db_url = f"sqlite:///{tmp_path / llm-usage.db}" + store = LLMUsageStore(db_url=db_url) + + store.record_usage( + task_type="summary", + provider_name="openai", + model_name="gpt-4o-mini", + prompt_tokens=120, + completion_tokens=80, + estimated_cost_usd=0.0001, + metadata={"estimated": True}, + ) + store.record_usage( + task_type="reasoning", + provider_name="anthropic", + model_name="claude-3-5-sonnet-20241022", + prompt_tokens=300, + completion_tokens=200, + estimated_cost_usd=0.001, + metadata={"estimated": True}, + ) + + summary = store.summarize(days=7) + + assert summary["totals"]["calls"] == 2 + assert summary["totals"]["total_tokens"] == 700 + assert len(summary["daily"]) >= 1 + assert len(summary["provider_models"]) == 2 + top = summary["provider_models"][0] + assert top["provider_name"] in {"openai", "anthropic"} diff --git a/tests/unit/test_model_endpoint_store.py b/tests/unit/test_model_endpoint_store.py new file mode 100644 index 00000000..e7315cef --- /dev/null +++ b/tests/unit/test_model_endpoint_store.py @@ -0,0 +1,103 @@ +from __future__ import annotations + +from pathlib import Path + +from paperbot.infrastructure.stores.model_endpoint_store import ModelEndpointStore + + +def test_activate_endpoint_switches_default(tmp_path: Path): + db_url = f"sqlite:///{tmp_path / store-activate.db}" + store = ModelEndpointStore(db_url=db_url) + + p1 = store.upsert_endpoint( + payload={ + "name": "p1", + "vendor": "openai_compatible", + "api_key_env": "OPENAI_API_KEY", + "models": ["gpt-4o-mini"], + "enabled": True, + "is_default": True, + } + ) + p2 = store.upsert_endpoint( + payload={ + "name": "p2", + "vendor": "openai_compatible", + "api_key_env": "OPENAI_API_KEY", + "models": ["gpt-4o"], + "enabled": True, + "is_default": False, + } + ) + + activated = store.activate_endpoint(int(p2["id"])) + assert activated is not None + assert activated["id"] == int(p2["id"]) + assert activated["is_default"] is True + + rows = store.list_endpoints() + row_by_name = {row["name"]: row for row in rows} + assert row_by_name["p1"]["is_default"] is False + assert row_by_name["p2"]["is_default"] is True + + +def test_delete_default_reassigns_new_default(tmp_path: Path): + db_url = f"sqlite:///{tmp_path / store-delete.db}" + store = ModelEndpointStore(db_url=db_url) + + p1 = store.upsert_endpoint( + payload={ + "name": "p1", + "vendor": "openai_compatible", + "api_key_env": "OPENAI_API_KEY", + "models": ["gpt-4o-mini"], + "enabled": True, + "is_default": True, + } + ) + p2 = store.upsert_endpoint( + payload={ + "name": "p2", + "vendor": "openai_compatible", + "api_key_env": "OPENAI_API_KEY", + "models": ["gpt-4o"], + "enabled": True, + "is_default": False, + } + ) + + assert store.delete_endpoint(int(p1["id"])) is True + + rows = store.list_endpoints() + assert len(rows) == 1 + assert rows[0]["id"] == int(p2["id"]) + assert rows[0]["is_default"] is True + + +def test_masked_api_key_write_back_keeps_existing_secret(tmp_path: Path): + db_url = f"sqlite:///{tmp_path / store-mask.db}" + store = ModelEndpointStore(db_url=db_url) + + created = store.upsert_endpoint( + payload={ + "name": "masked-provider", + "vendor": "openai_compatible", + "api_key_env": "OPENAI_API_KEY", + "api_key": "sk-live-1234567890", + "models": ["gpt-4o-mini"], + "enabled": True, + "is_default": True, + } + ) + endpoint_id = int(created["id"]) + assert created["api_key"].startswith("***") + + updated = store.upsert_endpoint( + payload={"api_key": created["api_key"]}, + endpoint_id=endpoint_id, + ) + assert updated["api_key"] == created["api_key"] + + raw = store.get_endpoint(endpoint_id, include_secrets=True) + assert raw is not None + assert raw["api_key"] == "sk-live-1234567890" diff --git a/tests/unit/test_model_endpoints_gateway.py b/tests/unit/test_model_endpoints_gateway.py new file mode 100644 index 00000000..823955e6 --- /dev/null +++ b/tests/unit/test_model_endpoints_gateway.py @@ -0,0 +1,192 @@ +from __future__ import annotations + +from pathlib import Path + +from fastapi.testclient import TestClient + +from paperbot.api import main as api_main +from paperbot.api.routes import model_endpoints as model_endpoints_route +from paperbot.infrastructure.llm.router import ModelRouter +from paperbot.infrastructure.stores.llm_usage_store import LLMUsageStore +from paperbot.infrastructure.stores.model_endpoint_store import ModelEndpointStore + + +def test_model_endpoint_crud_routes(tmp_path: Path, monkeypatch): + db_url = f"sqlite:///{tmp_path / model-endpoints.db}" + store = ModelEndpointStore(db_url=db_url) + monkeypatch.setattr(model_endpoints_route, "_store", store) + + with TestClient(api_main.app) as client: + created = client.post( + "/api/model-endpoints", + json={ + "name": "OpenRouter", + "vendor": "openai_compatible", + "base_url": "https://openrouter.ai/api/v1", + "api_key_env": "OPENROUTER_API_KEY", + "api_key": "sk-openrouter-123456789", + "models": ["deepseek/deepseek-r1"], + "task_types": ["reasoning", "summary"], + "enabled": True, + "is_default": True, + }, + ) + assert created.status_code == 200 + endpoint_id = created.json()["item"]["id"] + assert created.json()["item"]["api_key"].startswith("***") + + created2 = client.post( + "/api/model-endpoints", + json={ + "name": "AltProvider", + "vendor": "openai_compatible", + "base_url": "https://alt.example/v1", + "api_key_env": "ALT_API_KEY", + "api_key": "sk-alt-0987654321", + "models": ["alt/model-x"], + "task_types": ["default"], + "enabled": True, + "is_default": False, + }, + ) + assert created2.status_code == 200 + endpoint2_id = created2.json()["item"]["id"] + + listed = client.get("/api/model-endpoints") + assert listed.status_code == 200 + assert len(listed.json()["items"]) == 2 + + updated = client.patch( + f"/api/model-endpoints/{endpoint_id}", + json={ + "models": ["deepseek/deepseek-r1", "openai/gpt-4o-mini"], + "api_key": created.json()["item"]["api_key"], + }, + ) + assert updated.status_code == 200 + assert len(updated.json()["item"]["models"]) == 2 + + tested = client.post(f"/api/model-endpoints/{endpoint_id}/test", json={"remote": False}) + assert tested.status_code == 200 + assert tested.json()["ok"] is True + + activated = client.post(f"/api/model-endpoints/{endpoint2_id}/activate") + assert activated.status_code == 200 + assert activated.json()["item"]["id"] == endpoint2_id + assert activated.json()["item"]["is_default"] is True + + listed_after_activate = client.get("/api/model-endpoints") + assert listed_after_activate.status_code == 200 + items = listed_after_activate.json()["items"] + defaults = [item for item in items if item["is_default"]] + assert len(defaults) == 1 + assert defaults[0]["id"] == endpoint2_id + + deleted = client.delete(f"/api/model-endpoints/{endpoint_id}") + assert deleted.status_code == 200 + + deleted2 = client.delete(f"/api/model-endpoints/{endpoint2_id}") + assert deleted2.status_code == 200 + + listed_again = client.get("/api/model-endpoints") + assert listed_again.status_code == 200 + assert listed_again.json()["items"] == [] + + +def test_model_router_reads_registry_before_env(tmp_path: Path, monkeypatch): + db_url = f"sqlite:///{tmp_path / router-registry.db}" + store = ModelEndpointStore(db_url=db_url) + store.upsert_endpoint( + payload={ + "name": "GatewayDefault", + "vendor": "openai_compatible", + "base_url": "https://gateway.example/v1", + "api_key_env": "GATEWAY_API_KEY", + "models": ["gateway/model-a"], + "task_types": ["default", "summary"], + "enabled": True, + "is_default": True, + } + ) + + monkeypatch.setenv("PAPERBOT_DB_URL", db_url) + router = ModelRouter.from_env() + + models = router.list_models() + assert "GatewayDefault" in models + assert models["GatewayDefault"] == "openai:gateway/model-a" + + +def test_model_router_applies_task_routes_from_registry(tmp_path: Path, monkeypatch): + db_url = f"sqlite:///{tmp_path / router-routing.db}" + store = ModelEndpointStore(db_url=db_url) + store.upsert_endpoint( + payload={ + "name": "SummaryProvider", + "vendor": "openai_compatible", + "base_url": "https://summary.example/v1", + "api_key_env": "SUMMARY_API_KEY", + "models": ["summary/model"], + "task_types": ["summary"], + "enabled": True, + "is_default": False, + } + ) + store.upsert_endpoint( + payload={ + "name": "ReasoningProvider", + "vendor": "openai_compatible", + "base_url": "https://reasoning.example/v1", + "api_key_env": "REASONING_API_KEY", + "models": ["reasoning/model"], + "task_types": ["reasoning", "review"], + "enabled": True, + "is_default": False, + } + ) + store.upsert_endpoint( + payload={ + "name": "DefaultProvider", + "vendor": "openai_compatible", + "base_url": "https://default.example/v1", + "api_key_env": "DEFAULT_API_KEY", + "models": ["default/model"], + "task_types": ["chat", "default"], + "enabled": True, + "is_default": True, + } + ) + + monkeypatch.setenv("PAPERBOT_DB_URL", db_url) + router = ModelRouter.from_env() + + assert router._task_routing["summary"] == "SummaryProvider" + assert router._task_routing["reasoning"] == "ReasoningProvider" + assert router._task_routing["review"] == "ReasoningProvider" + assert router._task_routing["chat"] == "DefaultProvider" + assert router._task_routing["default"] == "DefaultProvider" + + +def test_llm_usage_summary_route(tmp_path: Path, monkeypatch): + db_url = f"sqlite:///{tmp_path / usage-route.db}" + store = ModelEndpointStore(db_url=db_url) + usage_store = LLMUsageStore(db_url=db_url) + monkeypatch.setattr(model_endpoints_route, "_store", store) + monkeypatch.setattr(model_endpoints_route, "_usage_store", usage_store) + + usage_store.record_usage( + task_type="summary", + provider_name="openai", + model_name="gpt-4o-mini", + prompt_tokens=100, + completion_tokens=40, + estimated_cost_usd=0.0, + ) + + with TestClient(api_main.app) as client: + resp = client.get("/api/model-endpoints/usage?days=7") + + assert resp.status_code == 200 + payload = resp.json()["summary"] + assert payload["totals"]["calls"] >= 1 + assert payload["totals"]["total_tokens"] >= 140 diff --git a/tests/unit/test_paper_judge_persistence.py b/tests/unit/test_paper_judge_persistence.py index 6e094991..cd80dca3 100644 --- a/tests/unit/test_paper_judge_persistence.py +++ b/tests/unit/test_paper_judge_persistence.py @@ -1,12 +1,15 @@ from __future__ import annotations +from datetime import datetime, timedelta, timezone from pathlib import Path from sqlalchemy import select +from paperbot.domain.identity import PaperIdentity from paperbot.infrastructure.stores.models import PaperJudgeScoreModel from paperbot.infrastructure.stores.paper_store import SqlAlchemyPaperStore from paperbot.infrastructure.stores.research_store import SqlAlchemyResearchStore +from paperbot.infrastructure.stores.identity_store import IdentityStore def _judged_report(): @@ -135,3 +138,104 @@ def test_saved_list_and_detail_from_research_store(tmp_path: Path): assert detail is not None assert detail["paper"]["title"] == "UniICL" assert detail["reading_status"]["status"] == "read" + + +def test_feedback_resolves_via_identity_store_mapping(tmp_path: Path): + db_path = tmp_path / "identity-feedback.db" + db_url = f"sqlite:///{db_path}" + + paper_store = SqlAlchemyPaperStore(db_url=db_url) + research_store = SqlAlchemyResearchStore(db_url=db_url) + identity_store = IdentityStore(db_url=db_url) + + paper = paper_store.upsert_paper( + paper={ + "title": "CrossSource Paper", + "url": "https://example.com/p/abc", + "pdf_url": "https://example.com/p/abc.pdf", + } + ) + + identity_store.upsert_identity( + paper_id=int(paper["id"]), + identity=PaperIdentity(source="papers_cool", external_id="pc:abc123"), + ) + + track = research_store.create_track(user_id="u2", name="track-u2", activate=True) + feedback = research_store.add_paper_feedback( + user_id="u2", + track_id=int(track["id"]), + paper_id="pc:abc123", + action="save", + metadata={}, + ) + + assert feedback is not None + assert feedback["paper_ref_id"] == int(paper["id"]) + + +def test_get_latest_judge_scores_returns_latest_row_per_paper(tmp_path: Path): + db_path = tmp_path / "judge-latest.db" + store = SqlAlchemyPaperStore(db_url=f"sqlite:///{db_path}") + + p1 = store.upsert_paper( + paper={ + "title": "P1", + "url": "https://example.com/p1", + "pdf_url": "https://example.com/p1.pdf", + } + ) + p2 = store.upsert_paper( + paper={ + "title": "P2", + "url": "https://example.com/p2", + "pdf_url": "https://example.com/p2.pdf", + } + ) + + now = datetime.now(timezone.utc) + with store._provider.session() as session: + session.add( + PaperJudgeScoreModel( + paper_id=int(p1["id"]), + query="q-old", + overall=3.1, + recommendation="worth_reading", + one_line_summary="old", + judge_model="m1", + judge_cost_tier=1, + scored_at=now - timedelta(days=1), + ) + ) + session.add( + PaperJudgeScoreModel( + paper_id=int(p1["id"]), + query="q-new", + overall=4.7, + recommendation="must_read", + one_line_summary="new", + judge_model="m2", + judge_cost_tier=1, + scored_at=now, + ) + ) + session.add( + PaperJudgeScoreModel( + paper_id=int(p2["id"]), + query="q-only", + overall=4.0, + recommendation="worth_reading", + one_line_summary="only", + judge_model="m3", + judge_cost_tier=1, + scored_at=now, + ) + ) + session.commit() + + latest = store.get_latest_judge_scores([int(p1["id"]), int(p2["id"]), -1]) + + assert set(latest.keys()) == {int(p1["id"]), int(p2["id"])} + assert latest[int(p1["id"])]["overall"] == 4.7 + assert latest[int(p1["id"])]["one_line_summary"] == "new" + assert latest[int(p2["id"])]["judge_model"] == "m3" diff --git a/tests/unit/test_paperscool_route.py b/tests/unit/test_paperscool_route.py index decf90d4..32a81218 100644 --- a/tests/unit/test_paperscool_route.py +++ b/tests/unit/test_paperscool_route.py @@ -2,6 +2,7 @@ from paperbot.api import main as api_main from paperbot.api.routes import paperscool as paperscool_route +from paperbot.infrastructure.stores.pipeline_session_store import PipelineSessionStore def _parse_sse_events(text: str): @@ -739,3 +740,190 @@ def _fake_enqueue(report): assert resp.status_code == 200 assert called["count"] == 1 + + +def test_paperscool_daily_resume_session(monkeypatch, tmp_path): + monkeypatch.setattr(paperscool_route, "PapersCoolTopicSearchWorkflow", _FakeWorkflow) + paperscool_route._pipeline_session_store = PipelineSessionStore( + db_url=f"sqlite:///{tmp_path / 'daily-session.db'}" + ) + + class _FakeJudgment: + def to_dict(self): + return { + "relevance": {"score": 5, "rationale": ""}, + "novelty": {"score": 4, "rationale": ""}, + "rigor": {"score": 4, "rationale": ""}, + "impact": {"score": 4, "rationale": ""}, + "clarity": {"score": 4, "rationale": ""}, + "overall": 4.2, + "one_line_summary": "good", + "recommendation": "must_read", + "judge_model": "fake", + "judge_cost_tier": 1, + } + + class _FakeJudge: + def __init__(self, llm_service=None): + pass + + def judge_single(self, *, paper, query): + return _FakeJudgment() + + def judge_with_calibration(self, *, paper, query, n_runs=1): + return _FakeJudgment() + + monkeypatch.setattr(paperscool_route, "PaperJudge", _FakeJudge) + monkeypatch.setattr(paperscool_route, "get_llm_service", lambda: object()) + + with TestClient(api_main.app) as client: + first = client.post( + "/api/research/paperscool/daily", + json={ + "queries": ["ICL压缩"], + "enable_judge": True, + "judge_runs": 1, + "judge_max_items_per_query": 2, + }, + ) + assert first.status_code == 200 + first_events = _parse_sse_events(first.text) + first_result = next(e for e in first_events if e.get("type") == "result") + session_id = first_result["data"].get("session_id") + assert session_id + + session_resp = client.get(f"/api/research/paperscool/sessions/{session_id}") + assert session_resp.status_code == 200 + assert session_resp.json()["session"]["status"] == "completed" + + resumed = client.post( + "/api/research/paperscool/daily", + json={ + "queries": ["ICL压缩"], + "enable_judge": True, + "session_id": session_id, + "resume": True, + }, + ) + assert resumed.status_code == 200 + resumed_events = _parse_sse_events(resumed.text) + resumed_result = next(e for e in resumed_events if e.get("type") == "result") + assert resumed_result["data"].get("resumed") is True + + +def test_paperscool_daily_route_pending_approval_and_queue(monkeypatch, tmp_path): + monkeypatch.setattr(paperscool_route, "PapersCoolTopicSearchWorkflow", _FakeWorkflow) + + calls = {"ingest": 0} + + def _fake_ingest(report): + calls["ingest"] += 1 + return {"saved": 1} + + monkeypatch.setattr(paperscool_route, "ingest_daily_report_to_registry", _fake_ingest) + monkeypatch.setattr( + paperscool_route, + "_pipeline_session_store", + PipelineSessionStore(db_url=f"sqlite:///{tmp_path / 'daily-approval.db'}"), + ) + + with TestClient(api_main.app) as client: + resp = client.post( + "/api/research/paperscool/daily", + json={ + "queries": ["ICL压缩"], + "require_approval": True, + }, + ) + assert resp.status_code == 200 + events = _parse_sse_events(resp.text) + types = [e.get("type") for e in events] + assert "approval_required" in types + result_event = next(e for e in events if e.get("type") == "result") + assert result_event["data"].get("approval_status") == "pending_approval" + session_id = result_event["data"].get("session_id") + assert session_id + + session_resp = client.get(f"/api/research/paperscool/sessions/{session_id}") + assert session_resp.status_code == 200 + assert session_resp.json()["session"]["status"] == "pending_approval" + + queue_resp = client.get("/api/research/paperscool/approvals?limit=10") + assert queue_resp.status_code == 200 + ids = [item["session_id"] for item in queue_resp.json().get("items", [])] + assert session_id in ids + + # Ingest is gated until explicit approve + assert calls["ingest"] == 0 + + +def test_paperscool_daily_approval_decisions(monkeypatch, tmp_path): + monkeypatch.setattr(paperscool_route, "PapersCoolTopicSearchWorkflow", _FakeWorkflow) + monkeypatch.setattr( + paperscool_route, + "_pipeline_session_store", + PipelineSessionStore(db_url=f"sqlite:///{tmp_path / 'daily-approval-decision.db'}"), + ) + monkeypatch.setattr( + paperscool_route, + "ingest_daily_report_to_registry", + lambda report: {"saved": len(report.get("queries") or [])}, + ) + monkeypatch.setattr( + paperscool_route, + "persist_judge_scores_to_registry", + lambda report: {"saved": 0}, + ) + + repo_calls = {"count": 0} + + def _fake_enqueue(_report): + repo_calls["count"] += 1 + + monkeypatch.setattr(paperscool_route, "_enqueue_repo_enrichment_async", _fake_enqueue) + + with TestClient(api_main.app) as client: + # Session A -> approve + pending_a = client.post( + "/api/research/paperscool/daily", + json={"queries": ["ICL压缩"], "require_approval": True}, + ) + assert pending_a.status_code == 200 + events_a = _parse_sse_events(pending_a.text) + result_a = next(e for e in events_a if e.get("type") == "result") + session_a = result_a["data"]["session_id"] + + approve = client.post(f"/api/research/paperscool/sessions/{session_a}/approve", json={}) + assert approve.status_code == 200 + approved_session = approve.json()["session"] + assert approved_session["status"] == "completed" + assert approved_session["result"]["approval_status"] == "approved" + assert "registry_ingest" in approved_session["result"]["report"] + + # Session B -> reject + pending_b = client.post( + "/api/research/paperscool/daily", + json={"queries": ["RAG"], "require_approval": True}, + ) + assert pending_b.status_code == 200 + events_b = _parse_sse_events(pending_b.text) + result_b = next(e for e in events_b if e.get("type") == "result") + session_b = result_b["data"]["session_id"] + + reject = client.post( + f"/api/research/paperscool/sessions/{session_b}/reject", + json={"reason": "Not ready"}, + ) + assert reject.status_code == 200 + rejected_session = reject.json()["session"] + assert rejected_session["status"] == "rejected" + assert rejected_session["state"].get("reject_reason") == "Not ready" + assert rejected_session["result"].get("approval_status") == "rejected" + + queue_resp = client.get("/api/research/paperscool/approvals?limit=10") + assert queue_resp.status_code == 200 + ids = [item["session_id"] for item in queue_resp.json().get("items", [])] + assert session_a not in ids + assert session_b not in ids + + assert repo_calls["count"] == 1 diff --git a/tests/unit/test_pipeline_session_store.py b/tests/unit/test_pipeline_session_store.py new file mode 100644 index 00000000..7c218d2e --- /dev/null +++ b/tests/unit/test_pipeline_session_store.py @@ -0,0 +1,69 @@ +from __future__ import annotations + +from pathlib import Path + +from paperbot.infrastructure.stores.pipeline_session_store import PipelineSessionStore + + +def test_pipeline_session_store_resume_and_result(tmp_path: Path): + db_url = f"sqlite:///{tmp_path / 'pipeline-session.db'}" + store = PipelineSessionStore(db_url=db_url) + + started = store.start_session( + workflow="paperscool_daily", + payload={"q": ["icl"]}, + session_id="sess-1", + resume=False, + ) + assert started["session_id"] == "sess-1" + assert started["status"] == "running" + + cp = store.save_checkpoint( + session_id="sess-1", + checkpoint="report_built", + state={"report": {"title": "Daily"}}, + ) + assert cp is not None + assert cp["checkpoint"] == "report_built" + + done = store.save_result( + session_id="sess-1", + status="completed", + result={"report": {"title": "Daily"}, "markdown": "# test"}, + ) + assert done is not None + assert done["status"] == "completed" + + resumed = store.start_session( + workflow="paperscool_daily", + payload={"q": ["icl"]}, + session_id="sess-1", + resume=True, + ) + assert resumed["status"] == "completed" + assert resumed["result"]["markdown"] == "# test" + + +def test_pipeline_session_store_list_and_update_status(tmp_path: Path): + db_url = f"sqlite:///{tmp_path / 'pipeline-session-list.db'}" + store = PipelineSessionStore(db_url=db_url) + + store.start_session(workflow="paperscool_daily", session_id="sess-a", payload={"q": ["a"]}) + store.start_session(workflow="paperscool_daily", session_id="sess-b", payload={"q": ["b"]}) + + updated = store.update_status( + session_id="sess-a", + status="pending_approval", + checkpoint="approval_pending", + state_patch={"flag": True}, + result={"approval_status": "pending_approval"}, + ) + assert updated is not None + assert updated["status"] == "pending_approval" + assert updated["checkpoint"] == "approval_pending" + assert updated["state"]["flag"] is True + assert updated["result"]["approval_status"] == "pending_approval" + + rows = store.list_sessions(workflow="paperscool_daily", status="pending_approval", limit=10) + assert len(rows) == 1 + assert rows[0]["session_id"] == "sess-a" diff --git a/tests/unit/test_research_paper_registry_routes.py b/tests/unit/test_research_paper_registry_routes.py index 5c65d7eb..c475dbcf 100644 --- a/tests/unit/test_research_paper_registry_routes.py +++ b/tests/unit/test_research_paper_registry_routes.py @@ -103,3 +103,105 @@ def test_paper_repos_route(tmp_path, monkeypatch): assert payload["paper_id"] == str(paper_id) assert len(payload["repos"]) == 1 assert payload["repos"][0]["full_name"] == "example/unicicl" + + +def test_track_feed_route_with_pagination_and_feedback_boost(tmp_path, monkeypatch): + db_path = tmp_path / "track-feed.db" + db_url = f"sqlite:///{db_path}" + + paper_store = SqlAlchemyPaperStore(db_url=db_url) + research_store = SqlAlchemyResearchStore(db_url=db_url) + + p1 = paper_store.upsert_paper( + paper={ + "title": "Retrieval-Augmented Generation in Practice", + "abstract": "rag retrieval pipeline", + "url": "https://example.com/p1", + } + ) + p2 = paper_store.upsert_paper( + paper={ + "title": "General Foundation Models", + "abstract": "broad model overview", + "url": "https://example.com/p2", + } + ) + paper_store.upsert_paper( + paper={ + "title": "Unrelated Database Benchmark", + "abstract": "oltp benchmark", + "url": "https://example.com/p3", + } + ) + + track = research_store.create_track( + user_id="u-feed", + name="rag-track", + keywords=["rag", "retrieval"], + activate=True, + ) + research_store.add_paper_feedback( + user_id="u-feed", + track_id=int(track["id"]), + paper_id=str(p2["id"]), + action="save", + metadata={"title": "General Foundation Models"}, + ) + + monkeypatch.setattr(research_route, "_research_store", research_store) + + with TestClient(api_main.app) as client: + page1 = client.get( + f"/api/research/tracks/{int(track[id])}/feed", + params={"user_id": "u-feed", "limit": 1, "offset": 0}, + ) + page2 = client.get( + f"/api/research/tracks/{int(track[id])}/feed", + params={"user_id": "u-feed", "limit": 1, "offset": 1}, + ) + + assert page1.status_code == 200 + assert page2.status_code == 200 + + payload1 = page1.json() + payload2 = page2.json() + + assert payload1["total"] >= 2 + assert len(payload1["items"]) == 1 + assert len(payload2["items"]) == 1 + assert payload1["items"][0]["paper"]["id"] != payload2["items"][0]["paper"]["id"] + + ids = {payload1["items"][0]["paper"]["id"], payload2["items"][0]["paper"]["id"]} + assert int(p1["id"]) in ids + assert int(p2["id"]) in ids + + +def test_deadline_radar_route_returns_workflow_query_and_track_match(tmp_path, monkeypatch): + db_path = tmp_path / "deadline-radar.db" + db_url = f"sqlite:///{db_path}" + research_store = SqlAlchemyResearchStore(db_url=db_url) + research_store.create_track( + user_id="u-deadline", + name="nlp-track", + keywords=["llm", "retrieval"], + activate=True, + ) + + monkeypatch.setattr(research_route, "_research_store", research_store) + + with TestClient(api_main.app) as client: + resp = client.get( + "/api/research/deadlines/radar", + params={"user_id": "u-deadline", "days": 365, "ccf_levels": "A"}, + ) + + assert resp.status_code == 200 + payload = resp.json() + assert payload["items"] + + first = payload["items"][0] + assert isinstance(first.get("workflow_query"), str) + assert first["workflow_query"] + + matched_any = any(item.get("matched_tracks") for item in payload["items"]) + assert matched_any diff --git a/tests/unit/test_streaming_envelope.py b/tests/unit/test_streaming_envelope.py index 3d498168..8de0f431 100644 --- a/tests/unit/test_streaming_envelope.py +++ b/tests/unit/test_streaming_envelope.py @@ -9,6 +9,7 @@ async def _simple_stream(): yield StreamEvent(type="progress", data={"phase": "judge", "message": "running"}) + yield StreamEvent(type="search_done", data={"ok": True}) yield StreamEvent(type="result", data={"ok": True}) @@ -28,7 +29,7 @@ async def test_wrap_generator_injects_envelope(): continue payloads.append(json.loads(data)) - assert len(payloads) == 2 + assert len(payloads) == 3 for idx, payload in enumerate(payloads, start=1): env = payload["envelope"] assert env["workflow"] == "paperscool_analyze" @@ -36,5 +37,9 @@ async def test_wrap_generator_injects_envelope(): assert env["trace_id"] == "trace_x" assert env["seq"] == idx assert isinstance(env["ts"], str) + assert env["event"] in {"progress", "result"} assert payloads[0]["envelope"]["phase"] == "judge" + assert payloads[0]["event"] == "progress" + assert payloads[1]["event"] == "progress" + assert payloads[2]["event"] == "result" diff --git a/web/src/app/api/model-endpoints/[id]/activate/route.ts b/web/src/app/api/model-endpoints/[id]/activate/route.ts new file mode 100644 index 00000000..4c64a5d6 --- /dev/null +++ b/web/src/app/api/model-endpoints/[id]/activate/route.ts @@ -0,0 +1,8 @@ +export const runtime = "nodejs" + +import { apiBaseUrl, proxyJson } from "../../../research/_base" + +export async function POST(req: Request, ctx: { params: Promise<{ id: string }> }) { + const { id } = await ctx.params + return proxyJson(req, `${apiBaseUrl()}/api/model-endpoints/${encodeURIComponent(id)}/activate`, "POST") +} diff --git a/web/src/app/api/model-endpoints/[id]/route.ts b/web/src/app/api/model-endpoints/[id]/route.ts new file mode 100644 index 00000000..c4041d00 --- /dev/null +++ b/web/src/app/api/model-endpoints/[id]/route.ts @@ -0,0 +1,13 @@ +export const runtime = "nodejs" + +import { apiBaseUrl, proxyJson } from "../../research/_base" + +export async function PATCH(req: Request, ctx: { params: Promise<{ id: string }> }) { + const { id } = await ctx.params + return proxyJson(req, `${apiBaseUrl()}/api/model-endpoints/${encodeURIComponent(id)}`, "PATCH") +} + +export async function DELETE(req: Request, ctx: { params: Promise<{ id: string }> }) { + const { id } = await ctx.params + return proxyJson(req, `${apiBaseUrl()}/api/model-endpoints/${encodeURIComponent(id)}`, "DELETE") +} diff --git a/web/src/app/api/model-endpoints/[id]/test/route.ts b/web/src/app/api/model-endpoints/[id]/test/route.ts new file mode 100644 index 00000000..20a02fdf --- /dev/null +++ b/web/src/app/api/model-endpoints/[id]/test/route.ts @@ -0,0 +1,8 @@ +export const runtime = "nodejs" + +import { apiBaseUrl, proxyJson } from "../../../research/_base" + +export async function POST(req: Request, ctx: { params: Promise<{ id: string }> }) { + const { id } = await ctx.params + return proxyJson(req, `${apiBaseUrl()}/api/model-endpoints/${encodeURIComponent(id)}/test`, "POST") +} diff --git a/web/src/app/api/model-endpoints/route.ts b/web/src/app/api/model-endpoints/route.ts new file mode 100644 index 00000000..b0b1bc06 --- /dev/null +++ b/web/src/app/api/model-endpoints/route.ts @@ -0,0 +1,12 @@ +export const runtime = "nodejs" + +import { apiBaseUrl, proxyJson } from "../research/_base" + +export async function GET(req: Request) { + const url = new URL(req.url) + return proxyJson(req, `${apiBaseUrl()}/api/model-endpoints?${url.searchParams.toString()}`, "GET") +} + +export async function POST(req: Request) { + return proxyJson(req, `${apiBaseUrl()}/api/model-endpoints`, "POST") +} diff --git a/web/src/app/api/research/deadlines/radar/route.ts b/web/src/app/api/research/deadlines/radar/route.ts new file mode 100644 index 00000000..515ac901 --- /dev/null +++ b/web/src/app/api/research/deadlines/radar/route.ts @@ -0,0 +1,8 @@ +export const runtime = "nodejs" + +import { apiBaseUrl, proxyJson } from "../../_base" + +export async function GET(req: Request) { + const url = new URL(req.url) + return proxyJson(req, `${apiBaseUrl()}/api/research/deadlines/radar?${url.searchParams.toString()}`, "GET") +} diff --git a/web/src/app/api/research/paperscool/approvals/route.ts b/web/src/app/api/research/paperscool/approvals/route.ts new file mode 100644 index 00000000..4a182cd7 --- /dev/null +++ b/web/src/app/api/research/paperscool/approvals/route.ts @@ -0,0 +1,12 @@ +export const runtime = "nodejs" + +import { apiBaseUrl, proxyJson } from "@/app/api/research/_base" + +export async function GET(req: Request) { + const url = new URL(req.url) + const query = url.searchParams.toString() + const upstream = query + ? `${apiBaseUrl()}/api/research/paperscool/approvals?${query}` + : `${apiBaseUrl()}/api/research/paperscool/approvals` + return proxyJson(req, upstream, "GET") +} diff --git a/web/src/app/api/research/paperscool/sessions/[sessionId]/approve/route.ts b/web/src/app/api/research/paperscool/sessions/[sessionId]/approve/route.ts new file mode 100644 index 00000000..0f0dba8e --- /dev/null +++ b/web/src/app/api/research/paperscool/sessions/[sessionId]/approve/route.ts @@ -0,0 +1,12 @@ +import { apiBaseUrl, proxyJson } from "@/app/api/research/_base" + +export const runtime = "nodejs" + +export async function POST(req: Request, ctx: { params: Promise<{ sessionId: string }> }) { + const { sessionId } = await ctx.params + return proxyJson( + req, + `${apiBaseUrl()}/api/research/paperscool/sessions/${encodeURIComponent(sessionId)}/approve`, + "POST", + ) +} diff --git a/web/src/app/api/research/paperscool/sessions/[sessionId]/reject/route.ts b/web/src/app/api/research/paperscool/sessions/[sessionId]/reject/route.ts new file mode 100644 index 00000000..cef3afb0 --- /dev/null +++ b/web/src/app/api/research/paperscool/sessions/[sessionId]/reject/route.ts @@ -0,0 +1,12 @@ +import { apiBaseUrl, proxyJson } from "@/app/api/research/_base" + +export const runtime = "nodejs" + +export async function POST(req: Request, ctx: { params: Promise<{ sessionId: string }> }) { + const { sessionId } = await ctx.params + return proxyJson( + req, + `${apiBaseUrl()}/api/research/paperscool/sessions/${encodeURIComponent(sessionId)}/reject`, + "POST", + ) +} diff --git a/web/src/app/api/research/paperscool/sessions/[sessionId]/route.ts b/web/src/app/api/research/paperscool/sessions/[sessionId]/route.ts new file mode 100644 index 00000000..a7c51356 --- /dev/null +++ b/web/src/app/api/research/paperscool/sessions/[sessionId]/route.ts @@ -0,0 +1,17 @@ +import { NextResponse } from "next/server" + +import { apiBaseUrl } from "@/app/api/research/_base" + +export async function GET(_req: Request, ctx: { params: Promise<{ sessionId: string }> }) { + const { sessionId } = await ctx.params + const upstream = await fetch( + `${apiBaseUrl()}/api/research/paperscool/sessions/${encodeURIComponent(sessionId)}`, + { cache: "no-store" }, + ) + + const body = await upstream.text() + return new NextResponse(body, { + status: upstream.status, + headers: { "Content-Type": upstream.headers.get("content-type") || "application/json" }, + }) +} diff --git a/web/src/app/api/research/tracks/[trackId]/feed/route.ts b/web/src/app/api/research/tracks/[trackId]/feed/route.ts new file mode 100644 index 00000000..328d0865 --- /dev/null +++ b/web/src/app/api/research/tracks/[trackId]/feed/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)}/feed?${url.searchParams.toString()}`, + "GET", + ) +} diff --git a/web/src/app/dashboard/page.tsx b/web/src/app/dashboard/page.tsx index e0310ca8..a52e4768 100644 --- a/web/src/app/dashboard/page.tsx +++ b/web/src/app/dashboard/page.tsx @@ -4,20 +4,22 @@ import { PipelineStatus } from "@/components/dashboard/PipelineStatus" import { ReadingQueue } from "@/components/dashboard/ReadingQueue" import { LLMUsageChart } from "@/components/dashboard/LLMUsageChart" import { QuickActions } from "@/components/dashboard/QuickActions" +import { DeadlineRadar } from "@/components/dashboard/DeadlineRadar" import { Users, FileText, Zap, BookOpen, Search, TrendingUp } from "lucide-react" -import { fetchStats, fetchTrendingTopics, fetchPipelineTasks, fetchReadingQueue, fetchLLMUsage } from "@/lib/api" +import { fetchStats, fetchTrendingTopics, fetchPipelineTasks, fetchReadingQueue, fetchLLMUsage, fetchDeadlineRadar } from "@/lib/api" import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card" import { Input } from "@/components/ui/input" import { Button } from "@/components/ui/button" import { Badge } from "@/components/ui/badge" export default async function DashboardPage() { - const [statsResult, trendsResult, tasksResult, readingQueueResult, llmUsageResult] = await Promise.allSettled([ + const [statsResult, trendsResult, tasksResult, readingQueueResult, llmUsageResult, deadlineResult] = await Promise.allSettled([ fetchStats(), fetchTrendingTopics(), fetchPipelineTasks(), fetchReadingQueue(), - fetchLLMUsage() + fetchLLMUsage(), + fetchDeadlineRadar("default"), ]) const stats = statsResult.status === "fulfilled" ? statsResult.value : { tracked_scholars: 0, @@ -28,7 +30,13 @@ export default async function DashboardPage() { const trends = trendsResult.status === "fulfilled" ? trendsResult.value : [] const tasks = tasksResult.status === "fulfilled" ? tasksResult.value : [] const readingQueue = readingQueueResult.status === "fulfilled" ? readingQueueResult.value : [] - const llmUsage = llmUsageResult.status === "fulfilled" ? llmUsageResult.value : [] + const usageSummary = llmUsageResult.status === "fulfilled" ? llmUsageResult.value : { + window_days: 7, + daily: [], + provider_models: [], + totals: { calls: 0, total_tokens: 0, total_cost_usd: 0 }, + } + const deadlines = deadlineResult.status === "fulfilled" ? deadlineResult.value : [] return (
@@ -53,7 +61,7 @@ export default async function DashboardPage() {
- +
@@ -68,12 +76,13 @@ export default async function DashboardPage() {
+
{/* Bottom Row */}
- +
@@ -83,16 +92,51 @@ export default async function DashboardPage() { -
- {trends.map((topic) => ( - - {topic.text} - - ))} +
+
+
+

Calls

+

{usageSummary.totals?.calls || 0}

+
+
+

Tokens

+

{(usageSummary.totals?.total_tokens || 0).toLocaleString()}

+
+
+

Cost (USD)

+

${Number(usageSummary.totals?.total_cost_usd || 0).toFixed(4)}

+
+
+ +
+ {(usageSummary.provider_models || []).slice(0, 6).map((row) => ( +
+
+

{row.provider_name} / {row.model_name}

+

calls: {row.calls}

+
+
+

{row.total_tokens.toLocaleString()} tok

+

${Number(row.total_cost_usd || 0).toFixed(4)}

+
+
+ ))} + {(!usageSummary.provider_models || usageSummary.provider_models.length === 0) && ( +
No usage records yet.
+ )} +
+ +
+ {trends.map((topic) => ( + + {topic.text} + + ))} +
diff --git a/web/src/app/page.tsx b/web/src/app/page.tsx index ef326328..0d4a4ac2 100644 --- a/web/src/app/page.tsx +++ b/web/src/app/page.tsx @@ -1,9 +1,5 @@ -import ResearchPageNew from "@/components/research/ResearchPageNew" +import { redirect } from "next/navigation" export default function HomePage() { - return ( -
- -
- ) + redirect("/dashboard") } diff --git a/web/src/app/research/page.tsx b/web/src/app/research/page.tsx index 800d3136..030b05c6 100644 --- a/web/src/app/research/page.tsx +++ b/web/src/app/research/page.tsx @@ -1,10 +1,13 @@ -import ResearchPageNew from "@/components/research/ResearchPageNew" +import { Suspense } from "react" + +import ResearchSplitWorkspace from "@/components/research/ResearchSplitWorkspace" export default function ResearchPage() { return (
- + Loading research workspace...
}> + +
) } - diff --git a/web/src/app/settings/page.tsx b/web/src/app/settings/page.tsx index 4899c0fe..e6748c00 100644 --- a/web/src/app/settings/page.tsx +++ b/web/src/app/settings/page.tsx @@ -1,44 +1,468 @@ +"use client" + +import { useEffect, useMemo, useState } from "react" +import { CheckCircle2, Loader2, PlugZap, Plus, Save, Trash2, Wrench } from "lucide-react" + +import { Badge } from "@/components/ui/badge" import { Button } from "@/components/ui/button" -import { Input } from "@/components/ui/input" import { Card, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card" +import { Input } from "@/components/ui/input" + +type ModelEndpoint = { + id: number + name: string + vendor: string + base_url?: string | null + api_key_env: string + api_key?: string + models: string[] + task_types: string[] + enabled: boolean + is_default: boolean + api_key_present?: boolean +} + +type ModelEndpointListResponse = { + items: ModelEndpoint[] +} + +type FormState = { + id?: number + name: string + vendor: string + base_url: string + api_key_env: string + api_key: string + models: string + task_types: string + enabled: boolean + is_default: boolean +} + +type Preset = { + label: string + name: string + vendor: string + base_url: string + api_key_env: string + models: string[] + task_types: string[] +} + +const QUICK_PRESETS: Preset[] = [ + { + label: "OpenAI", + name: "OpenAI", + vendor: "openai", + base_url: "https://api.openai.com/v1", + api_key_env: "OPENAI_API_KEY", + models: ["gpt-4o-mini"], + task_types: ["default", "summary", "chat"], + }, + { + label: "OpenRouter", + name: "OpenRouter", + vendor: "openai_compatible", + base_url: "https://openrouter.ai/api/v1", + api_key_env: "OPENROUTER_API_KEY", + models: ["openai/gpt-4o-mini"], + task_types: ["reasoning", "review"], + }, + { + label: "Anthropic", + name: "Anthropic", + vendor: "anthropic", + base_url: "", + api_key_env: "ANTHROPIC_API_KEY", + models: ["claude-3-5-sonnet-20241022"], + task_types: ["reasoning", "analysis"], + }, + { + label: "Ollama", + name: "Local Ollama", + vendor: "ollama", + base_url: "http://localhost:11434", + api_key_env: "OLLAMA_API_KEY", + models: ["llama3.1"], + task_types: ["default", "chat"], + }, +] + +const EMPTY_FORM: FormState = { + name: "", + vendor: "openai_compatible", + base_url: "", + api_key_env: "OPENAI_API_KEY", + api_key: "", + models: "", + task_types: "", + enabled: true, + is_default: false, +} + +function toPayload(form: FormState) { + return { + name: form.name.trim(), + vendor: form.vendor, + base_url: form.base_url.trim() || null, + api_key_env: form.api_key_env.trim(), + api_key: form.api_key, + models: form.models + .split(",") + .map((x) => x.trim()) + .filter(Boolean), + task_types: form.task_types + .split(",") + .map((x) => x.trim()) + .filter(Boolean), + enabled: form.enabled, + is_default: form.is_default, + } +} export default function SettingsPage() { - return ( -
-

Settings

- -
- - - LLM Configuration - Configure your AI model providers and keys. - - -
- - -
-
- - -
- -
-
- - - - Notifications - Manage your alert preferences. - - -
- - -
-
-
-
-
- ) + const [items, setItems] = useState([]) + const [loading, setLoading] = useState(false) + const [saving, setSaving] = useState(false) + const [testingId, setTestingId] = useState(null) + const [activatingId, setActivatingId] = useState(null) + const [form, setForm] = useState(EMPTY_FORM) + const [error, setError] = useState(null) + const [message, setMessage] = useState(null) + + const editing = useMemo(() => typeof form.id === "number", [form.id]) + + const load = async () => { + setLoading(true) + setError(null) + try { + const res = await fetch("/api/model-endpoints") + if (!res.ok) throw new Error(`${res.status} ${res.statusText}`) + const payload = (await res.json()) as ModelEndpointListResponse + setItems(payload.items || []) + } catch (e) { + setError(e instanceof Error ? e.message : String(e)) + setItems([]) + } finally { + setLoading(false) + } + } + + useEffect(() => { + load().catch(() => {}) + }, []) + + function resetForm() { + setForm(EMPTY_FORM) + setMessage(null) + setError(null) + } + + function applyPreset(preset: Preset) { + setForm((prev) => ({ + ...prev, + name: preset.name, + vendor: preset.vendor, + base_url: preset.base_url, + api_key_env: preset.api_key_env, + models: preset.models.join(", "), + task_types: preset.task_types.join(", "), + enabled: true, + })) + setMessage(`Preset applied: ${preset.label}`) + setError(null) + } + + function editItem(item: ModelEndpoint) { + setForm({ + id: item.id, + name: item.name, + vendor: item.vendor, + base_url: item.base_url || "", + api_key_env: item.api_key_env, + api_key: item.api_key || "", + models: (item.models || []).join(", "), + task_types: (item.task_types || []).join(", "), + enabled: item.enabled, + is_default: item.is_default, + }) + setMessage(null) + setError(null) + } + + async function saveItem() { + setSaving(true) + setError(null) + setMessage(null) + try { + const payload = toPayload(form) + if (!payload.name) throw new Error("Name is required") + if (!payload.models.length) throw new Error("At least one model is required") + + const res = await fetch(editing ? `/api/model-endpoints/${form.id}` : "/api/model-endpoints", { + method: editing ? "PATCH" : "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(payload), + }) + if (!res.ok) { + const text = await res.text().catch(() => "") + throw new Error(text || `${res.status} ${res.statusText}`) + } + await load() + resetForm() + setMessage(editing ? "Provider updated." : "Provider created.") + } catch (e) { + setError(e instanceof Error ? e.message : String(e)) + } finally { + setSaving(false) + } + } + + async function deleteItem(id: number) { + if (!confirm("Delete this provider?")) return + setError(null) + setMessage(null) + try { + const res = await fetch(`/api/model-endpoints/${id}`, { method: "DELETE" }) + if (!res.ok) throw new Error(`${res.status} ${res.statusText}`) + await load() + if (form.id === id) { + resetForm() + } + setMessage("Provider removed.") + } catch (e) { + setError(e instanceof Error ? e.message : String(e)) + } + } + + async function activateItem(id: number) { + setActivatingId(id) + setError(null) + setMessage(null) + try { + const res = await fetch(`/api/model-endpoints/${id}/activate`, { method: "POST" }) + if (!res.ok) { + const text = await res.text().catch(() => "") + throw new Error(text || `${res.status} ${res.statusText}`) + } + await load() + setMessage("Provider activated.") + } catch (e) { + setError(e instanceof Error ? e.message : String(e)) + } finally { + setActivatingId(null) + } + } + + async function testItem(id: number) { + setTestingId(id) + setError(null) + setMessage(null) + try { + const res = await fetch(`/api/model-endpoints/${id}/test`, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ remote: false }), + }) + const payload = await res.json().catch(() => ({})) + if (!res.ok) { + throw new Error(String(payload?.detail || `${res.status} ${res.statusText}`)) + } + setMessage(payload?.message || "Connection test passed.") + } catch (e) { + setError(e instanceof Error ? e.message : String(e)) + } finally { + setTestingId(null) + } + } + + return ( +
+

Settings

+ + + + Model Providers + + Add providers, switch default route, test connectivity, and keep API keys masked in UI. + + + + {error &&

{error}

} + {message &&

{message}

} + +
+

Quick Presets

+
+ {QUICK_PRESETS.map((preset) => ( + + ))} +
+
+ +
+
+ + setForm((p) => ({ ...p, name: e.target.value }))} placeholder="DeepSeek via OpenRouter" /> +
+ +
+ + +
+ +
+ + setForm((p) => ({ ...p, base_url: e.target.value }))} + placeholder="https://openrouter.ai/api/v1" + /> +
+ +
+ + setForm((p) => ({ ...p, api_key_env: e.target.value }))} + placeholder="OPENAI_API_KEY" + /> +
+ +
+ + setForm((p) => ({ ...p, api_key: e.target.value }))} + placeholder={editing ? "***masked value" : "sk-..."} + /> +
+ +
+ + setForm((p) => ({ ...p, models: e.target.value }))} placeholder="gpt-4o-mini, gpt-4o" /> +
+ +
+ + setForm((p) => ({ ...p, task_types: e.target.value }))} + placeholder="default, summary, reasoning, code" + /> +
+ +
+ setForm((p) => ({ ...p, enabled: e.target.checked }))} + /> + +
+
+ setForm((p) => ({ ...p, is_default: e.target.checked }))} + /> + +
+
+ +
+ + +
+
+
+ + + + Configured Providers + Used by LLM service router for task-level model selection. + + + {loading ? ( +

Loading providers...

+ ) : !items.length ? ( +

No providers yet.

+ ) : ( + items.map((item) => ( +
+
+
+ {item.name} + {item.vendor} + {item.is_default && default} + {!item.enabled && disabled} + + {item.api_key_present ? "key ready" : "missing key"} + +
+ +
+ {!item.is_default && ( + + )} + + + +
+
+ +
+
base_url: {item.base_url || "(default)"}
+
api_key_env: {item.api_key_env}
+
api_key: {item.api_key || "(from env only)"}
+
models: {(item.models || []).join(", ") || "-"}
+
task_routes: {(item.task_types || []).join(", ") || "-"}
+
+
+ )) + )} +
+
+
+ ) } diff --git a/web/src/app/workflows/page.tsx b/web/src/app/workflows/page.tsx index 0affe008..df7fd76f 100644 --- a/web/src/app/workflows/page.tsx +++ b/web/src/app/workflows/page.tsx @@ -1,9 +1,23 @@ import TopicWorkflowDashboard from "@/components/research/TopicWorkflowDashboard" -export default function WorkflowsPage() { +type WorkflowsPageProps = { + searchParams?: Promise> +} + +export default async function WorkflowsPage({ searchParams }: WorkflowsPageProps) { + const params = searchParams ? await searchParams : {} + const rawQuery = params?.query + const queryValue = Array.isArray(rawQuery) ? rawQuery[0] : rawQuery + const initialQueries = queryValue + ? queryValue + .split(",") + .map((q) => q.trim()) + .filter(Boolean) + : undefined + return (
- +
) } diff --git a/web/src/components/ai-elements/code-block.tsx b/web/src/components/ai-elements/code-block.tsx new file mode 100644 index 00000000..b6f4b724 --- /dev/null +++ b/web/src/components/ai-elements/code-block.tsx @@ -0,0 +1,53 @@ +"use client" + +import { useState } from "react" +import { Check, Copy } from "lucide-react" + +import { Button } from "@/components/ui/button" +import { cn } from "@/lib/utils" + +interface CodeBlockProps { + code: string + title?: string + language?: string + className?: string +} + +export function CodeBlock({ code, title = "Code", language, className }: CodeBlockProps) { + const [copied, setCopied] = useState(false) + + const copy = async () => { + try { + await navigator.clipboard.writeText(code) + setCopied(true) + setTimeout(() => setCopied(false), 1200) + } catch { + setCopied(false) + } + } + + return ( +
+
+
+ {title} + {language ? ` · ${language}` : ""} +
+ +
+
+        {code}
+      
+
+ ) +} diff --git a/web/src/components/ai-elements/index.ts b/web/src/components/ai-elements/index.ts new file mode 100644 index 00000000..98028530 --- /dev/null +++ b/web/src/components/ai-elements/index.ts @@ -0,0 +1,3 @@ +export { CodeBlock } from "./code-block" +export { ReasoningBlock } from "./reasoning-block" +export { ToolActionsGroup } from "./tool-actions-group" diff --git a/web/src/components/ai-elements/reasoning-block.tsx b/web/src/components/ai-elements/reasoning-block.tsx new file mode 100644 index 00000000..6cd09267 --- /dev/null +++ b/web/src/components/ai-elements/reasoning-block.tsx @@ -0,0 +1,44 @@ +import { Lightbulb } from "lucide-react" + +import { Badge } from "@/components/ui/badge" +import { cn } from "@/lib/utils" + +interface ReasoningBlockProps { + reasons: string[] + title?: string + compact?: boolean + className?: string +} + +export function ReasoningBlock({ + reasons, + title = "Reasoning", + compact = false, + className, +}: ReasoningBlockProps) { + if (!reasons.length) return null + + return ( +
+
+ + {title} +
+ {compact ? ( +
+ {reasons.map((reason) => ( + + {reason} + + ))} +
+ ) : ( +
    + {reasons.map((reason) => ( +
  • {reason}
  • + ))} +
+ )} +
+ ) +} diff --git a/web/src/components/ai-elements/tool-actions-group.tsx b/web/src/components/ai-elements/tool-actions-group.tsx new file mode 100644 index 00000000..913bf845 --- /dev/null +++ b/web/src/components/ai-elements/tool-actions-group.tsx @@ -0,0 +1,48 @@ +"use client" + +import { ReactNode } from "react" + +import { Button } from "@/components/ui/button" +import { cn } from "@/lib/utils" + +type ToolAction = { + id: string + label: string + icon?: ReactNode + onClick?: () => void + disabled?: boolean + variant?: "default" | "outline" | "ghost" | "destructive" | "secondary" + className?: string +} + +interface ToolActionsGroupProps { + actions: ToolAction[] + className?: string + ariaLabel?: string +} + +export function ToolActionsGroup({ + actions, + className, + ariaLabel = "Tool actions", +}: ToolActionsGroupProps) { + if (!actions.length) return null + + return ( +
+ {actions.map((action) => ( + + ))} +
+ ) +} diff --git a/web/src/components/dashboard/DeadlineRadar.tsx b/web/src/components/dashboard/DeadlineRadar.tsx new file mode 100644 index 00000000..d12e951e --- /dev/null +++ b/web/src/components/dashboard/DeadlineRadar.tsx @@ -0,0 +1,79 @@ +import Link from "next/link" +import { CalendarClock, ExternalLink } from "lucide-react" + +import { Badge } from "@/components/ui/badge" +import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card" +import type { DeadlineRadarItem } from "@/lib/types" + +interface DeadlineRadarProps { + items: DeadlineRadarItem[] +} + +export function DeadlineRadar({ items }: DeadlineRadarProps) { + return ( + + + + + Deadline Radar + + + + {!items.length ? ( +

No upcoming deadlines in selected window.

+ ) : ( + items.map((item) => ( +
+
+
+

{item.name}

+

+ {item.field} · D-{item.days_left} +

+
+ + CCF {item.ccf_level} + +
+ + {!!item.matched_tracks?.length && ( +
+ {item.matched_tracks.slice(0, 2).map((track) => ( + + + Track: {track.track_name} + + + ))} +
+ )} + +
+ + Open in Workflows + + {item.url && ( + + Official CFP + + )} +
+
+ )) + )} +
+
+ ) +} diff --git a/web/src/components/dashboard/LLMUsageChart.tsx b/web/src/components/dashboard/LLMUsageChart.tsx index 81463923..11bbb06c 100644 --- a/web/src/components/dashboard/LLMUsageChart.tsx +++ b/web/src/components/dashboard/LLMUsageChart.tsx @@ -1,33 +1,78 @@ "use client" +import { useMemo } from "react" +import { Bar, BarChart, CartesianGrid, Legend, ResponsiveContainer, Tooltip, XAxis, YAxis } from "recharts" + import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card" -import { BarChart, Bar, XAxis, YAxis, CartesianGrid, Tooltip, ResponsiveContainer, Legend } from "recharts" -import type { LLMUsageRecord } from "@/lib/types" +import type { LLMUsageSummary } from "@/lib/types" interface LLMUsageChartProps { - data: LLMUsageRecord[] + data: LLMUsageSummary +} + +const BAR_COLORS = ["#10b981", "#8b5cf6", "#6b7280", "#3b82f6", "#f59e0b"] + +function formatProviderLabel(value: string): string { + if (!value) return "Unknown" + return value + .split(/[_-]/g) + .filter(Boolean) + .map((part) => part.charAt(0).toUpperCase() + part.slice(1)) + .join(" ") } export function LLMUsageChart({ data }: LLMUsageChartProps) { - return ( - - - LLM Token Usage (7 Days) - - - - - - - `${v / 1000}k`} /> - value !== undefined ? `${Number(value).toLocaleString()} tokens` : ''} /> - - - - - - - - - ) + const providerKeys = useMemo(() => { + const totals = new Map() + for (const row of data.daily || []) { + for (const [provider, count] of Object.entries(row.providers || {})) { + totals.set(provider, (totals.get(provider) || 0) + Number(count || 0)) + } + } + return [...totals.entries()] + .sort((a, b) => b[1] - a[1]) + .slice(0, 5) + .map(([name]) => name) + }, [data.daily]) + + const chartRows = useMemo(() => { + return (data.daily || []).map((row) => { + const next: Record = { + date: row.date, + total_tokens: row.total_tokens, + } + for (const key of providerKeys) { + next[key] = Number(row.providers?.[key] || 0) + } + return next + }) + }, [data.daily, providerKeys]) + + return ( + + + LLM Token Usage ({data.window_days} Days) + + + + + + + `${Math.round(Number(v) / 1000)}k`} /> + `${Number(value).toLocaleString()} tokens`} /> + formatProviderLabel(value)} /> + {providerKeys.map((provider, index) => ( + + ))} + + + + + ) } diff --git a/web/src/components/layout/Sidebar.tsx b/web/src/components/layout/Sidebar.tsx index 7585c5fa..6860783b 100644 --- a/web/src/components/layout/Sidebar.tsx +++ b/web/src/components/layout/Sidebar.tsx @@ -2,7 +2,6 @@ import Link from "next/link" import { usePathname } from "next/navigation" -import { useState } from "react" import { cn } from "@/lib/utils" import { Button } from "@/components/ui/button" import { @@ -24,8 +23,8 @@ type SidebarProps = React.HTMLAttributes & { } const routes = [ - { label: "Research", icon: FlaskConical, href: "/" }, { label: "Dashboard", icon: LayoutDashboard, href: "/dashboard" }, + { label: "Research", icon: FlaskConical, href: "/research" }, { label: "Scholars", icon: Users, href: "/scholars" }, { label: "Papers", icon: FileText, href: "/papers" }, { label: "Workflows", icon: Workflow, href: "/workflows" }, @@ -61,8 +60,7 @@ export function Sidebar({ className, collapsed, onToggle }: SidebarProps) { {/* Nav items */}
{routes.map((route) => { - const isActive = - route.href === "/" ? pathname === "/" : pathname.startsWith(route.href) + const isActive = pathname === route.href || pathname.startsWith(`${route.href}/`) return ( + + +
+
+ {mobileActive === "rail" && rail} + {mobileActive === "list" && list} + {mobileActive === "detail" && detail} +
+
+ ) + } + + return ( +
+
+ + + +
+ { + const normalized = { + rail: Number(next.rail || DEFAULT_LAYOUT.rail), + list: Number(next.list || DEFAULT_LAYOUT.list), + detail: Number(next.detail || DEFAULT_LAYOUT.detail), + } + setLayout(normalized) + try { + window.localStorage.setItem(layoutKey, JSON.stringify(normalized)) + } catch { + // Ignore storage failures. + } + }} + className="flex-1 min-h-0" + > + setCollapsedState("rail", inPixels < 2)} + className="min-h-0 overflow-auto" + > + {rail} + + + + + setCollapsedState("list", inPixels < 2)} + className="min-h-0 overflow-auto" + > + {list} + + + + + setCollapsedState("detail", inPixels < 2)} + className="min-h-0 overflow-auto" + > + {detail} + + +
+ ) +} diff --git a/web/src/components/research/ApprovalQueuePanel.tsx b/web/src/components/research/ApprovalQueuePanel.tsx new file mode 100644 index 00000000..5e683a30 --- /dev/null +++ b/web/src/components/research/ApprovalQueuePanel.tsx @@ -0,0 +1,122 @@ +"use client" + +import { useCallback, useEffect, useState } from "react" +import { Check, Loader2, RefreshCw, X } from "lucide-react" + +import { Button } from "@/components/ui/button" +import { Card, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card" + +type ApprovalItem = { + session_id: string + status: string + checkpoint?: string + updated_at?: string | null + title?: string + query_count?: number + unique_items?: number +} + +type ApprovalQueueResponse = { + items: ApprovalItem[] +} + +export function ApprovalQueuePanel() { + const [items, setItems] = useState([]) + const [loading, setLoading] = useState(false) + const [actingId, setActingId] = useState(null) + const [error, setError] = useState(null) + + const load = useCallback(async () => { + setLoading(true) + setError(null) + try { + const res = await fetch("/api/research/paperscool/approvals?limit=20", { cache: "no-store" }) + if (!res.ok) throw new Error(`${res.status} ${res.statusText}`) + const payload = (await res.json()) as ApprovalQueueResponse + setItems(payload.items || []) + } catch (e) { + setError(e instanceof Error ? e.message : String(e)) + setItems([]) + } finally { + setLoading(false) + } + }, []) + + useEffect(() => { + load().catch(() => {}) + }, [load]) + + async function decide(sessionId: string, action: "approve" | "reject") { + setActingId(sessionId) + setError(null) + try { + const body = action === "reject" ? { reason: "Rejected from UI queue" } : {} + const res = await fetch(`/api/research/paperscool/sessions/${encodeURIComponent(sessionId)}/${action}`, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(body), + }) + if (!res.ok) { + const text = await res.text().catch(() => "") + throw new Error(text || `${res.status} ${res.statusText}`) + } + await load() + } catch (e) { + setError(e instanceof Error ? e.message : String(e)) + } finally { + setActingId(null) + } + } + + return ( + + + Pending Approvals + Manual approve/reject queue for gated enrichment sessions. + + +
+ +
+ + {error &&

{error}

} + + {!items.length ? ( +

No pending approvals.

+ ) : ( + items.map((item) => ( +
+

{item.title || "DailyPaper Session"}

+

session: {item.session_id}

+

queries: {item.query_count || 0} · unique: {item.unique_items || 0}

+
+ + +
+
+ )) + )} +
+
+ ) +} diff --git a/web/src/components/research/FeedTab.tsx b/web/src/components/research/FeedTab.tsx new file mode 100644 index 00000000..9b42cdc5 --- /dev/null +++ b/web/src/components/research/FeedTab.tsx @@ -0,0 +1,136 @@ +"use client" + +import { useEffect, useMemo, useState } from "react" +import { Loader2, RefreshCw } from "lucide-react" + +import { Button } from "@/components/ui/button" +import { Card, CardContent } from "@/components/ui/card" + +import { PaperCard, type Paper } from "./PaperCard" + +type FeedItem = { + paper: { + id?: number + title?: string + abstract?: string + authors?: string[] + year?: number + venue?: string + citation_count?: number + url?: string + } + latest_judge?: Paper["latest_judge"] + latest_feedback_action?: string | null +} + +type FeedResponse = { + items: FeedItem[] + total: number + limit: number + offset: number +} + +interface FeedTabProps { + userId: string + trackId: number | null + onLike?: (paperId: string, rank: number) => Promise | void + onSave?: (paperId: string, rank: number, paper: Paper) => Promise | void + onDislike?: (paperId: string, rank: number) => Promise | void +} + +function toPaper(item: FeedItem): Paper { + const id = String(item.paper.id || "") + return { + paper_id: id, + title: item.paper.title || "Untitled", + abstract: item.paper.abstract || "", + authors: item.paper.authors || [], + year: item.paper.year, + venue: item.paper.venue, + citation_count: item.paper.citation_count || 0, + url: item.paper.url, + latest_judge: item.latest_judge, + is_saved: (item.latest_feedback_action || "").toLowerCase() === "save", + } +} + +export function FeedTab({ userId, trackId, onLike, onSave, onDislike }: FeedTabProps) { + const [items, setItems] = useState([]) + const [loading, setLoading] = useState(false) + const [error, setError] = useState(null) + + const papers = useMemo(() => items.map(toPaper), [items]) + + const load = async () => { + if (!trackId) { + setItems([]) + return + } + setLoading(true) + 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()}`) + if (!res.ok) { + throw new Error(`${res.status} ${res.statusText}`) + } + const payload = (await res.json()) as FeedResponse + setItems(payload.items || []) + } catch (e) { + setError(e instanceof Error ? e.message : String(e)) + setItems([]) + } finally { + setLoading(false) + } + } + + useEffect(() => { + load().catch(() => {}) + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [userId, trackId]) + + if (!trackId) { + return
Select a track to view feed.
+ } + + return ( +
+
+

Track feed from DailyPaper + personalized ranking.

+ +
+ + {error && ( + + {error} + + )} + + {loading && !papers.length ? ( +
Loading feed...
+ ) : !papers.length ? ( +
No feed items yet for this track.
+ ) : ( +
+ {papers.map((paper, idx) => ( + onLike(paper.paper_id, idx) : undefined} + onSave={onSave ? () => onSave(paper.paper_id, idx, paper) : undefined} + onDislike={onDislike ? () => onDislike(paper.paper_id, idx) : undefined} + /> + ))} +
+ )} +
+ ) +} diff --git a/web/src/components/research/MemoryTab.tsx b/web/src/components/research/MemoryTab.tsx new file mode 100644 index 00000000..4757d51b --- /dev/null +++ b/web/src/components/research/MemoryTab.tsx @@ -0,0 +1,119 @@ +"use client" + +import { useEffect, useState } from "react" +import { Loader2, RefreshCw } from "lucide-react" + +import { Button } from "@/components/ui/button" +import { Card, CardContent } from "@/components/ui/card" +import { Badge } from "@/components/ui/badge" + +type MemoryItem = { + id: number + kind?: string + content?: string + tags?: string[] + created_at?: string | null +} + +type MemoryResponse = { + user_id: string + items: MemoryItem[] +} + +interface MemoryTabProps { + userId: string + trackId: number | null +} + +export function MemoryTab({ userId, trackId }: MemoryTabProps) { + const [items, setItems] = useState([]) + const [loading, setLoading] = useState(false) + const [error, setError] = useState(null) + + const load = async () => { + if (!trackId) { + setItems([]) + return + } + setLoading(true) + 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()}`) + if (!res.ok) { + throw new Error(`${res.status} ${res.statusText}`) + } + const payload = (await res.json()) as MemoryResponse + setItems(payload.items || []) + } catch (e) { + setError(e instanceof Error ? e.message : String(e)) + setItems([]) + } finally { + setLoading(false) + } + } + + useEffect(() => { + load().catch(() => {}) + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [userId, trackId]) + + if (!trackId) { + return
Select a track to view memory items.
+ } + + return ( +
+
+

Memory inbox for current track.

+ +
+ + {error && ( + + {error} + + )} + + {loading && !items.length ? ( +
Loading memory items...
+ ) : !items.length ? ( +
No memory items for this track yet.
+ ) : ( +
+ {items.map((item) => ( + + +
+ + {item.kind || "note"} + + {item.created_at && ( + {new Date(item.created_at).toLocaleString()} + )} +
+

{item.content || "(empty)"}

+ {!!item.tags?.length && ( +
+ {item.tags.map((tag) => ( + + {tag} + + ))} +
+ )} +
+
+ ))} +
+ )} +
+ ) +} diff --git a/web/src/components/research/PaperCard.tsx b/web/src/components/research/PaperCard.tsx index da775353..0e84a7cc 100644 --- a/web/src/components/research/PaperCard.tsx +++ b/web/src/components/research/PaperCard.tsx @@ -1,11 +1,11 @@ "use client" -import { useState } from "react" +import { useEffect, useState } from "react" import { Check, ExternalLink, Heart, Loader2, Save, ThumbsDown } from "lucide-react" import { cn, safeHref } from "@/lib/utils" +import { ReasoningBlock, ToolActionsGroup } from "@/components/ai-elements" import { Badge } from "@/components/ui/badge" -import { Button } from "@/components/ui/button" export type Paper = { paper_id: string @@ -16,6 +16,13 @@ export type Paper = { citation_count?: number authors?: string[] url?: string + latest_judge?: { + overall?: number + recommendation?: string + one_line_summary?: string + judge_model?: string + } + is_saved?: boolean } interface PaperCardProps { @@ -39,7 +46,11 @@ export function PaperCard({ isLoading = false, className, }: PaperCardProps) { - const [isSaved, setIsSaved] = useState(false) + const [isSaved, setIsSaved] = useState(Boolean(paper.is_saved)) + + useEffect(() => { + setIsSaved(Boolean(paper.is_saved)) + }, [paper.is_saved]) const [isLiked, setIsLiked] = useState(false) const [isDisliked, setIsDisliked] = useState(false) const [actionLoading, setActionLoading] = useState(null) @@ -47,6 +58,9 @@ export function PaperCard({ const authorText = paper.authors?.slice(0, 3).join(", ") || "Unknown authors" const hasMoreAuthors = (paper.authors?.length || 0) > 3 const safeUrl = safeHref(paper.url) + const judge = paper.latest_judge + const judgeOverall = Number(judge?.overall || 0) + const judgeRec = String(judge?.recommendation || "").replace(/_/g, " ") const handleSave = async () => { if (!onSave || isSaved) return @@ -148,79 +162,98 @@ export function PaperCard({ {/* Recommendation reasons */} {reasons && reasons.length > 0 && ( + + )} + + {judge && judgeOverall > 0 && (
- {reasons.map((reason) => ( - - {reason} + + Judge {judgeOverall.toFixed(1)} + + {judgeRec && ( + + {judgeRec} - ))} + )}
)} {/* Action buttons */} -
- {onSave && ( - - )} - {onLike && ( - - )} - {onDislike && ( - - )} -
+ + ) : isSaved ? ( + + ) : ( + + ), + }, + ] + : []), + ...(onLike + ? [ + { + id: "like", + label: isLiked ? "Liked" : "Like", + variant: "ghost" as const, + className: cn("transition-all", isLiked && "text-red-500 hover:text-red-600"), + onClick: handleLike, + disabled: isLoading || actionLoading !== null, + icon: + actionLoading === "like" ? ( + + ) : ( + + ), + }, + ] + : []), + ...(onDislike + ? [ + { + id: "dislike", + label: isDisliked ? "Hidden" : "Not relevant", + variant: "ghost" as const, + className: cn( + "transition-all", + isDisliked + ? "text-orange-500 hover:text-orange-600" + : "text-muted-foreground hover:text-destructive" + ), + onClick: handleDislike, + disabled: isLoading || actionLoading !== null, + icon: + actionLoading === "dislike" ? ( + + ) : ( + + ), + }, + ] + : []), + ]} + /> ) } diff --git a/web/src/components/research/ResearchPageNew.tsx b/web/src/components/research/ResearchPageNew.tsx index 4022c97c..c45253a7 100644 --- a/web/src/components/research/ResearchPageNew.tsx +++ b/web/src/components/research/ResearchPageNew.tsx @@ -1,6 +1,7 @@ "use client" import { useEffect, useMemo, useState } from "react" +import { useSearchParams } from "next/navigation" import { cn } from "@/lib/utils" import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card" @@ -13,10 +14,14 @@ import { DialogTitle, } from "@/components/ui/dialog" import { Button } from "@/components/ui/button" +import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs" import { SearchBox } from "./SearchBox" import { TrackPills } from "./TrackPills" import { SearchResults } from "./SearchResults" +import { FeedTab } from "./FeedTab" +import { SavedTab } from "./SavedTab" +import { MemoryTab } from "./MemoryTab" import { CreateTrackModal } from "./CreateTrackModal" import { EditTrackModal } from "./EditTrackModal" import { ManageTracksModal } from "./ManageTracksModal" @@ -52,6 +57,8 @@ function getGreeting(): string { } export default function ResearchPageNew() { + const searchParams = useSearchParams() + // User state const [userId] = useState("default") @@ -64,6 +71,8 @@ export default function ResearchPageNew() { const [hasSearched, setHasSearched] = useState(false) const [isSearching, setIsSearching] = useState(false) const [contextPack, setContextPack] = useState(null) + const [activeTab, setActiveTab] = useState("search") + const [searchSources, setSearchSources] = useState(["semantic_scholar"]) // UI state const [loading, setLoading] = useState(false) @@ -85,6 +94,7 @@ export default function ResearchPageNew() { const papers = contextPack?.paper_recommendations || [] const reasons = contextPack?.paper_recommendation_reasons || {} + const routeTrackId = Number(searchParams.get("track_id") || 0) // Load tracks on mount useEffect(() => { @@ -92,6 +102,14 @@ export default function ResearchPageNew() { // eslint-disable-next-line react-hooks/exhaustive-deps }, []) + useEffect(() => { + if (!routeTrackId || !Number.isFinite(routeTrackId)) return + if (!tracks.some((track) => track.id === routeTrackId)) return + if (activeTrackId === routeTrackId) return + activateTrack(routeTrackId).catch(() => {}) + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [routeTrackId, tracks, activeTrackId]) + async function refreshTracks(): Promise { const data = await fetchJson<{ tracks: Track[] }>( `/api/research/tracks?user_id=${encodeURIComponent(userId)}` @@ -135,6 +153,7 @@ export default function ResearchPageNew() { query, paper_limit: 10, memory_limit: 8, + sources: searchSources, offline: false, include_cross_track: false, stage: "auto", @@ -150,6 +169,7 @@ export default function ResearchPageNew() { ) setContextPack(data.context_pack) + setActiveTab("search") } catch (e) { setError(String(e)) } finally { @@ -157,6 +177,17 @@ export default function ResearchPageNew() { } } + function toggleSearchSource(source: string) { + setSearchSources((prev) => { + const exists = prev.includes(source) + if (exists) { + const next = prev.filter((x) => x !== source) + return next.length ? next : ["semantic_scholar"] + } + return [...prev, source] + }) + } + async function handleCreateTrack(data: { name: string description: string @@ -314,14 +345,7 @@ export default function ResearchPageNew() { const trackToClearName = tracks.find((t) => t.id === trackToClear)?.name || "this track" return ( -
+
{/* Confirm Clear Memory Dialog */} @@ -398,7 +422,7 @@ export default function ResearchPageNew() {
{/* Greeting - only show before search */} @@ -460,16 +484,46 @@ export default function ResearchPageNew() { )} - {/* Search Results */} - handleFeedback(paperId, "like", rank)} - onSave={(paperId, rank, paper) => handleFeedback(paperId, "save", rank, paper)} - onDislike={(paperId, rank) => handleFeedback(paperId, "dislike", rank)} - /> + + + Search + Feed + Saved + Memory + + + + handleFeedback(paperId, "like", rank)} + onSave={(paperId, rank, paper) => handleFeedback(paperId, "save", rank, paper)} + onDislike={(paperId, rank) => handleFeedback(paperId, "dislike", rank)} + /> + + + + handleFeedback(paperId, "like", rank)} + onSave={(paperId, rank, paper) => handleFeedback(paperId, "save", rank, paper)} + onDislike={(paperId, rank) => handleFeedback(paperId, "dislike", rank)} + /> + + + + + + + + + +
) diff --git a/web/src/components/research/ResearchSplitWorkspace.tsx b/web/src/components/research/ResearchSplitWorkspace.tsx new file mode 100644 index 00000000..d00669e1 --- /dev/null +++ b/web/src/components/research/ResearchSplitWorkspace.tsx @@ -0,0 +1,81 @@ +"use client" + +import Link from "next/link" +import { BookOpen, Library, LayoutPanelLeft, SearchCheck, Settings2 } from "lucide-react" + +import { SplitPanels } from "@/components/layout/SplitPanels" +import { ApprovalQueuePanel } from "@/components/research/ApprovalQueuePanel" +import ResearchPageNew from "@/components/research/ResearchPageNew" +import SavedPapersList from "@/components/research/SavedPapersList" +import { Badge } from "@/components/ui/badge" +import { Button } from "@/components/ui/button" +import { Card, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card" + +function RailPanel() { + return ( +
+ + + + Research Rail + + Track-scoped shortcuts and workspace entry points. + + + Search + Feed + Memory +
+ + + +
+
+
+
+ ) +} + +function DetailPanel() { + return ( +
+ + + + + Saved Snapshot + + Live view from papers library for quick context. + + + +
+ ) +} + +export default function ResearchSplitWorkspace() { + return ( + } + list={} + detail={} + className="h-[calc(100vh-4rem)]" + /> + ) +} diff --git a/web/src/components/research/SavedTab.tsx b/web/src/components/research/SavedTab.tsx new file mode 100644 index 00000000..0a526df1 --- /dev/null +++ b/web/src/components/research/SavedTab.tsx @@ -0,0 +1,118 @@ +"use client" + +import { useEffect, useMemo, useState } from "react" +import { Loader2, RefreshCw } from "lucide-react" + +import { Button } from "@/components/ui/button" +import { Card, CardContent } from "@/components/ui/card" + +import { PaperCard, type Paper } from "./PaperCard" + +type SavedItem = { + paper: { + id?: number + title?: string + abstract?: string + authors?: string[] + year?: number + venue?: string + citation_count?: number + url?: string + } + latest_judge?: Paper["latest_judge"] + saved_at?: string | null +} + +type SavedResponse = { + user_id: string + items: SavedItem[] +} + +interface SavedTabProps { + userId: string + trackId: number | null +} + +function toPaper(item: SavedItem): Paper { + return { + paper_id: String(item.paper.id || ""), + title: item.paper.title || "Untitled", + abstract: item.paper.abstract || "", + authors: item.paper.authors || [], + year: item.paper.year, + venue: item.paper.venue, + citation_count: item.paper.citation_count || 0, + url: item.paper.url, + latest_judge: item.latest_judge, + is_saved: true, + } +} + +export function SavedTab({ userId, trackId }: SavedTabProps) { + const [items, setItems] = useState([]) + const [loading, setLoading] = useState(false) + const [error, setError] = useState(null) + + const papers = useMemo(() => items.map(toPaper), [items]) + + const load = async () => { + setLoading(true) + setError(null) + try { + const qs = new URLSearchParams({ + user_id: userId, + sort_by: "saved_at", + limit: "100", + }) + if (trackId) { + qs.set("track_id", String(trackId)) + } + const res = await fetch(`/api/research/papers/saved?${qs.toString()}`) + if (!res.ok) { + throw new Error(`${res.status} ${res.statusText}`) + } + const payload = (await res.json()) as SavedResponse + setItems(payload.items || []) + } catch (e) { + setError(e instanceof Error ? e.message : String(e)) + setItems([]) + } finally { + setLoading(false) + } + } + + useEffect(() => { + load().catch(() => {}) + // eslint-disable-next-line react-hooks/exhaustive-deps + }, [userId, trackId]) + + return ( +
+
+

Saved papers scoped to current track.

+ +
+ + {error && ( + + {error} + + )} + + {loading && !papers.length ? ( +
Loading saved papers...
+ ) : !papers.length ? ( +
No saved papers in this track.
+ ) : ( +
+ {papers.map((paper, idx) => ( + + ))} +
+ )} +
+ ) +} diff --git a/web/src/components/research/SearchResults.tsx b/web/src/components/research/SearchResults.tsx index 4d2a4a28..d5a408c7 100644 --- a/web/src/components/research/SearchResults.tsx +++ b/web/src/components/research/SearchResults.tsx @@ -13,11 +13,19 @@ interface SearchResultsProps { reasons?: Record isSearching?: boolean hasSearched: boolean + selectedSources?: string[] + onToggleSource?: (source: string) => void onLike?: (paperId: string, rank: number) => Promise | void onSave?: (paperId: string, rank: number, paper: Paper) => Promise | void onDislike?: (paperId: string, rank: number) => Promise | void } +const SOURCE_OPTIONS: Array<{ value: string; label: string }> = [ + { value: "semantic_scholar", label: "S2" }, + { value: "arxiv", label: "arXiv" }, + { value: "openalex", label: "OpenAlex" }, +] + function PaperCardSkeleton() { return (
@@ -46,6 +54,8 @@ export function SearchResults({ reasons, isSearching = false, hasSearched, + selectedSources = ["semantic_scholar"], + onToggleSource, onLike, onSave, onDislike, @@ -93,7 +103,26 @@ export function SearchResults({

Found {papers.length} papers

- {/* Future: Add sort dropdown here */} +
+ {SOURCE_OPTIONS.map((source) => { + const active = selectedSources.includes(source.value) + return ( + + ) + })} +
diff --git a/web/src/components/research/TopicWorkflowDashboard.tsx b/web/src/components/research/TopicWorkflowDashboard.tsx index 5a2de3c2..e67a989b 100644 --- a/web/src/components/research/TopicWorkflowDashboard.tsx +++ b/web/src/components/research/TopicWorkflowDashboard.tsx @@ -689,9 +689,15 @@ function NewsletterSubscribeWidget() { /* ── Main Dashboard ───────────────────────────────────── */ -export default function TopicWorkflowDashboard() { +type TopicWorkflowDashboardProps = { + initialQueries?: string[] +} + +export default function TopicWorkflowDashboard({ initialQueries }: TopicWorkflowDashboardProps = {}) { /* Config state (local — queries only) */ - const [queryItems, setQueryItems] = useState([...DEFAULT_QUERIES]) + const [queryItems, setQueryItems] = useState([ + ...((initialQueries && initialQueries.length ? initialQueries : DEFAULT_QUERIES) || DEFAULT_QUERIES), + ]) /* Persisted state (zustand) */ const store = useWorkflowStore() diff --git a/web/src/components/studio/ExecutionLog.tsx b/web/src/components/studio/ExecutionLog.tsx index 36cb780c..a547729b 100644 --- a/web/src/components/studio/ExecutionLog.tsx +++ b/web/src/components/studio/ExecutionLog.tsx @@ -2,6 +2,7 @@ import { useState } from "react" import { useStudioStore, AgentAction } from "@/lib/store/studio-store" +import { CodeBlock } from "@/components/ai-elements" import { ScrollArea } from "@/components/ui/scroll-area" import { Bot, FileCode, Wrench, Plug, AlertCircle, CheckCircle2, Search, Terminal, ChevronDown, ChevronRight, Clock, Sparkles } from "lucide-react" import { cn } from "@/lib/utils" @@ -50,6 +51,8 @@ function ActionItem({ action, onViewDiff, isLast }: ActionItemProps) { const colors = actionColors[iconKey] || actionColors[action.type] || actionColors.text const hasExpandableContent = Boolean(action.metadata?.params || action.metadata?.result || action.metadata?.mcpResult) + const stringifyPayload = (payload: unknown): string => + typeof payload === "string" ? payload : JSON.stringify(payload, null, 2) || "" return (
@@ -107,22 +110,10 @@ function ActionItem({ action, onViewDiff, isLast }: ActionItemProps) { {expanded && (
{Boolean(action.metadata.params) && ( -
-
Args
-
-                                                    {JSON.stringify(action.metadata.params, null, 2)}
-                                                
-
+ )} {Boolean(action.metadata.result) && ( -
-
Result
-
-                                                    {typeof action.metadata.result === 'string'
-                                                        ? action.metadata.result
-                                                        : JSON.stringify(action.metadata.result, null, 2)}
-                                                
-
+ )}
)} @@ -144,13 +135,11 @@ function ActionItem({ action, onViewDiff, isLast }: ActionItemProps) { )}
{expanded && Boolean(action.metadata.mcpResult) && ( -
-
-                                            {typeof action.metadata.mcpResult === 'string'
-                                                ? action.metadata.mcpResult
-                                                : JSON.stringify(action.metadata.mcpResult, null, 2)}
-                                        
-
+ )}
) : action.type === 'error' ? ( diff --git a/web/src/lib/api.ts b/web/src/lib/api.ts index ac433540..30683720 100644 --- a/web/src/lib/api.ts +++ b/web/src/lib/api.ts @@ -1,4 +1,17 @@ -import { Activity, Paper, PaperDetails, Scholar, ScholarDetails, Stats, WikiConcept, TrendingTopic, PipelineTask, ReadingQueueItem, LLMUsageRecord } from "./types" +import { + Activity, + Paper, + PaperDetails, + Scholar, + ScholarDetails, + Stats, + WikiConcept, + TrendingTopic, + PipelineTask, + ReadingQueueItem, + LLMUsageSummary, + DeadlineRadarItem, +} from "./types" const API_BASE_URL = process.env.PAPERBOT_API_BASE_URL || "http://127.0.0.1:8000/api" @@ -26,9 +39,20 @@ async function postJson(path: string, payload: Record): Prom } export async function fetchStats(): Promise { - // TODO: Replace with real API call - // const res = await fetch(`${API_BASE_URL}/stats`) - // return res.json() + try { + const usage = await fetchLLMUsage() + const tokenCount = usage.totals.total_tokens + const prettyTokens = tokenCount >= 1000 ? `${Math.round(tokenCount / 1000)}k` : `${tokenCount}` + return { + tracked_scholars: 128, + new_papers: 12, + llm_usage: prettyTokens, + read_later: 8, + } + } catch { + // Keep resilient dashboard fallback. + } + return { tracked_scholars: 128, new_papers: 12, @@ -112,16 +136,75 @@ export async function fetchReadingQueue(): Promise { ] } -export async function fetchLLMUsage(): Promise { - return [ - { date: "Mon", gpt4: 12000, claude: 8000, ollama: 3000 }, - { date: "Tue", gpt4: 15000, claude: 9500, ollama: 4000 }, - { date: "Wed", gpt4: 10000, claude: 7000, ollama: 5000 }, - { date: "Thu", gpt4: 18000, claude: 12000, ollama: 2000 }, - { date: "Fri", gpt4: 14000, claude: 10000, ollama: 6000 }, - { date: "Sat", gpt4: 8000, claude: 5000, ollama: 1000 }, - { date: "Sun", gpt4: 6000, claude: 4000, ollama: 500 } - ] +export async function fetchLLMUsage(days: number = 7): Promise { + try { + const qs = new URLSearchParams({ days: String(days) }) + const res = await fetch(`${API_BASE_URL}/model-endpoints/usage?${qs.toString()}`, { + cache: "no-store", + }) + if (!res.ok) throw new Error("usage endpoint unavailable") + const payload = await res.json() as { summary?: LLMUsageSummary } + if (payload.summary) { + return payload.summary + } + } catch { + // Keep static fallback for local-first UX. + } + + return { + window_days: days, + daily: [ + { + date: "Mon", + total_tokens: 23000, + total_cost_usd: 0.0, + providers: { openai: 12000, anthropic: 8000, ollama: 3000 }, + }, + { + date: "Tue", + total_tokens: 28500, + total_cost_usd: 0.0, + providers: { openai: 15000, anthropic: 9500, ollama: 4000 }, + }, + { + date: "Wed", + total_tokens: 22000, + total_cost_usd: 0.0, + providers: { openai: 10000, anthropic: 7000, ollama: 5000 }, + }, + { + date: "Thu", + total_tokens: 32000, + total_cost_usd: 0.0, + providers: { openai: 18000, anthropic: 12000, ollama: 2000 }, + }, + ], + provider_models: [], + totals: { + calls: 0, + total_tokens: 105500, + total_cost_usd: 0, + }, + } +} + +export async function fetchDeadlineRadar(userId: string = "default"): Promise { + try { + const qs = new URLSearchParams({ + user_id: userId, + days: "180", + ccf_levels: "A,B,C", + limit: "10", + }) + const res = await fetch(`${API_BASE_URL}/research/deadlines/radar?${qs.toString()}`, { + cache: "no-store", + }) + if (!res.ok) return [] + const payload = await res.json() as { items?: DeadlineRadarItem[] } + return payload.items || [] + } catch { + return [] + } } diff --git a/web/src/lib/sse.ts b/web/src/lib/sse.ts index 9b5ecc29..451bea6e 100644 --- a/web/src/lib/sse.ts +++ b/web/src/lib/sse.ts @@ -4,11 +4,13 @@ export type StreamEnvelope = { trace_id?: string seq?: number phase?: string | null + event?: string ts?: string } export type SSEMessage = { type?: string + event?: string data?: unknown message?: string | null envelope?: StreamEnvelope | null @@ -16,11 +18,39 @@ export type SSEMessage = { export type NormalizedSSEEvent = { type: string + event: "status" | "progress" | "tool" | "result" | "error" | "done" data: unknown message: string | null envelope: StreamEnvelope } +const PROGRESS_LIKE = new Set([ + "progress", + "search_done", + "report_built", + "llm_summary", + "llm_done", + "trend", + "insight", + "judge", + "judge_done", + "filter_done", +]) + +function normalizeEventKind(message: SSEMessage): "status" | "progress" | "tool" | "result" | "error" | "done" { + const direct = typeof message.event === "string" ? message.event : undefined + const fromEnvelope = typeof message.envelope?.event === "string" ? message.envelope.event : undefined + const raw = (direct || fromEnvelope || message.type || "").toLowerCase() + + if (raw === "status") return "status" + if (raw === "result") return "result" + if (raw === "error") return "error" + if (raw === "done") return "done" + if (raw.startsWith("tool")) return "tool" + if (PROGRESS_LIKE.has(raw)) return "progress" + return "status" +} + function asEnvelope(raw: unknown): StreamEnvelope { if (!raw || typeof raw !== "object") return {} const obj = raw as Record @@ -30,6 +60,7 @@ function asEnvelope(raw: unknown): StreamEnvelope { trace_id: typeof obj.trace_id === "string" ? obj.trace_id : undefined, seq: typeof obj.seq === "number" ? obj.seq : undefined, phase: typeof obj.phase === "string" ? obj.phase : null, + event: typeof obj.event === "string" ? obj.event : undefined, ts: typeof obj.ts === "string" ? obj.ts : undefined, } } @@ -41,6 +72,7 @@ export function normalizeSSEMessage(message: SSEMessage, fallbackWorkflow = "unk return { type: typeof message.type === "string" && message.type.length > 0 ? message.type : "unknown", + event: normalizeEventKind(message), data: message.data, message: typeof message.message === "string" ? message.message : null, envelope: { @@ -49,6 +81,7 @@ export function normalizeSSEMessage(message: SSEMessage, fallbackWorkflow = "unk trace_id: envelope.trace_id, seq: envelope.seq, phase: derivedPhase, + event: envelope.event, ts: envelope.ts, }, } diff --git a/web/src/lib/types.ts b/web/src/lib/types.ts index 0096d23e..d66ba219 100644 --- a/web/src/lib/types.ts +++ b/web/src/lib/types.ts @@ -123,9 +123,44 @@ export interface ReadingQueueItem { priority: number } -export interface LLMUsageRecord { +export interface LLMUsageDailyRecord { date: string - gpt4: number - claude: number - ollama: number + total_tokens: number + total_cost_usd: number + providers: Record +} + +export interface LLMUsageProviderModelRecord { + provider_name: string + model_name: string + calls: number + total_tokens: number + total_cost_usd: number +} + +export interface LLMUsageSummary { + window_days: number + daily: LLMUsageDailyRecord[] + provider_models: LLMUsageProviderModelRecord[] + totals: { + calls: number + total_tokens: number + total_cost_usd: number + } +} + +export interface DeadlineRadarItem { + name: string + ccf_level: string + field: string + deadline: string + days_left: number + url: string + keywords: string[] + workflow_query: string + matched_tracks: Array<{ + track_id: number + track_name: string + matched_keywords: string[] + }> }