Spaces:
Runtime error
Runtime error
File size: 5,601 Bytes
03037be | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 | import sqlite3
import faiss
import time
import json
from sentence_transformers import SentenceTransformer
from fastapi import FastAPI, HTTPException, Query
from contextlib import asynccontextmanager
from loguru import logger
from functools import lru_cache
from src.api.cache import get_cached_response, set_cache
from src.rag.pipeline import summarize_paper
from src.api.models import (
SearchResponse, PaperResult, PaperDetail, HealthResponse
)
# config
DB_PATH = "data/papers.db"
INDEX_PATH = "data/faiss_index.bin"
EMBEDDINGS_PATH = "data/embeddings.npy"
ID_MAP_PATH = "data/id_map.json"
MODEL_NAME = "all-MiniLM-L6-v2"
NPROBE = 10
# global state
state = {
"index": None,
"id_map": None,
"model": None,
"db_conn": None,
"n_papers": 0,
}
@asynccontextmanager
async def lifespan(app: FastAPI):
logger.info("Starting up API...")
logger.info("Loading FAISS index...")
state["index"] = faiss.read_index(INDEX_PATH)
state["index"].nprobe = NPROBE
logger.info(f"FAISS index loaded with {state['index'].ntotal} vectors.")
logger.info("Loading sentence transformer model...")
state["model"] = SentenceTransformer(MODEL_NAME)
logger.success("Sentence transformer model loaded.")
logger.info("Loading ID map...")
with open(ID_MAP_PATH, "r") as f:
state["id_map"] = json.load(f)
logger.success(f"ID map loaded β {len(state['id_map'])} entries β")
logger.info("Connecting to SQLite database...")
state["db_conn"] = sqlite3.connect(DB_PATH, check_same_thread=False)
cursor = state["db_conn"].cursor()
cursor.execute("SELECT COUNT(*) FROM papers")
state["n_papers"] = cursor.fetchone()[0]
logger.success(f"Connected β {state['n_papers']} papers β")
logger.info("API startup complete.")
yield
logger.info("Shutting down API...")
if state["db_conn"]:
state["db_conn"].close()
logger.info("API shutdown complete.")
app = FastAPI(
title="PaperLens API",
description="Semantic search engine for research papers",
version="1.0",
lifespan=lifespan,
)
# ββ helpers βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
def get_paper_by_id(paper_id: str) -> dict | None:
cursor = state["db_conn"].cursor()
cursor.execute(
"SELECT id, title, authors, abstract, year, field, url FROM papers WHERE id = ?",
(paper_id,)
)
row = cursor.fetchone()
if row:
return {
"id": row[0],
"title": row[1],
"authors": row[2],
"abstract": row[3],
"year": row[4],
"field": row[5],
"url": row[6],
}
return None
# ββ endpoints βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
@app.get("/health", response_model=HealthResponse)
def health():
return HealthResponse(
status="ok",
paper_indexed=state["n_papers"],
index_loading=state["index"] is not None,
embedding_loading=state["model"] is not None,
)
@lru_cache(maxsize=512)
def embed_query(query: str):
return state["model"].encode(
[query],
normalize_embeddings=True,
convert_to_numpy=True,
).astype("float32")
@app.get("/search", response_model=SearchResponse)
def search(
query: str = Query(..., description="Search query"),
top_k: int = Query(10, ge=1, le=50, description="Number of results"),
):
if not query.strip():
raise HTTPException(status_code=400, detail="Query cannot be empty")
cached = get_cached_response(query, top_k)
if cached:
return cached
start = time.time()
query_vector = embed_query(query)
distances, indices = state["index"].search(query_vector, top_k)
results = []
for idx, dist in zip(indices[0], distances[0]):
if idx == -1:
continue
paper_id = state["id_map"].get(str(idx))
if not paper_id:
continue
paper = get_paper_by_id(paper_id)
if not paper:
continue
results.append(PaperResult(**paper, score=float(dist)))
latency_ms = (time.time() - start) * 1000
logger.info(f"Search: {len(results)} results in {latency_ms:.2f}ms")
response = SearchResponse(
query=query,
results=results,
total=len(results),
latency_ms=round(latency_ms, 3),
)
set_cache(query, top_k, response) # FIX: moved before return
return response
# FIX: was indented inside search() β now a top-level route
@app.get("/papers/{paper_id}", response_model=PaperDetail)
def get_paper(paper_id: str):
paper = get_paper_by_id(paper_id)
if not paper:
raise HTTPException(status_code=404, detail="Paper not found")
return PaperDetail(**paper)
# FIX: now returns a dict with summary + title instead of raw string
@app.get("/summarize/{paper_id}")
def summarize(paper_id: str):
paper = get_paper_by_id(paper_id)
if not paper:
raise HTTPException(status_code=404, detail=f"Paper {paper_id} not found")
summary = summarize_paper(
title=paper["title"],
abstract=paper["abstract"],
)
return {
"paper_id": paper_id,
"title": paper["title"],
"summary": summary,
} |