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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
78 changes: 75 additions & 3 deletions backend/app/repositories/search_repository.py
Original file line number Diff line number Diff line change
@@ -1,21 +1,37 @@
from typing import List, Tuple
from datetime import date
from typing import List, Optional, Tuple, Union

from sqlalchemy import select
from sqlalchemy import ColumnElement, and_, or_, select
from sqlalchemy.ext.asyncio import AsyncSession

from app.models.paper import Paper as PaperModel
from app.schemas.search_dto import (
AdvancedSearchFilter,
ConditionGroup,
TextCondition,
)

# Mapping from DTO field names to SQLAlchemy model columns
_FIELD_COLUMN = {
"title": PaperModel.title,
"abstract": PaperModel.abstract,
}


class SearchRepository:
"""Repository for search-related database operations."""

@staticmethod
async def search_papers_by_embeddings(
db: AsyncSession, embeddings: List[List[float]], limit: int = 5
db: AsyncSession,
embeddings: List[List[float]],
limit: int = 5,
search_filter: Optional[AdvancedSearchFilter] = None,
) -> List[Tuple[PaperModel, float]]:
"""
Perform a vector search for papers based on a list of embeddings.
Returns a list of (PaperModel, avg_distance) tuples ordered by ascending distance.
Optionally applies advanced search filters (year range, text conditions).
"""

# Build distance expressions
Expand All @@ -29,6 +45,62 @@ async def search_papers_by_embeddings(
.limit(limit)
)

if search_filter:
clauses = SearchRepository._build_filter_clauses(search_filter)
if clauses:
stmt = stmt.where(and_(*clauses))

result = await db.execute(stmt)
rows = result.fetchall()
return [(paper, float(dist)) for paper, dist in rows]

@staticmethod
def _build_filter_clauses(search_filter: AdvancedSearchFilter) -> list:
"""Build a list of top-level SQLAlchemy filter clauses from the advanced filter."""
clauses = []

if search_filter.year_from is not None:
clauses.append(PaperModel.published_at >= date(search_filter.year_from, 1, 1))

if search_filter.year_to is not None:
clauses.append(PaperModel.published_at <= date(search_filter.year_to, 12, 31))

condition_clause = SearchRepository._build_group_clause(search_filter.root)
if condition_clause is not None:
clauses.append(condition_clause)

return clauses

@staticmethod
def _build_group_clause(group: ConditionGroup) -> Optional[ColumnElement[bool]]:
"""Recursively build an AND/OR clause from a ConditionGroup."""
if not group.children:
return None

child_clauses = []
for child in group.children:
clause = SearchRepository._build_node_clause(child)
if clause is not None:
child_clauses.append(clause)

if not child_clauses:
return None

if group.operator == "AND":
return and_(*child_clauses)
return or_(*child_clauses)

@staticmethod
def _build_node_clause(
node: Union[TextCondition, ConditionGroup]
) -> Optional[ColumnElement[bool]]:
"""Build a clause for a single node (condition or nested group)."""
if node.type == "group":
return SearchRepository._build_group_clause(node)

column = _FIELD_COLUMN[node.field]
pattern = f"%{node.value}%"

if node.operator == "contains":
return column.ilike(pattern)
return or_(column.is_(None), ~column.ilike(pattern))
50 changes: 37 additions & 13 deletions backend/app/routes/search_routes.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,10 @@
import json

from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile, status
from sqlalchemy.ext.asyncio import AsyncSession

from app.core.database import get_db
from app.schemas.search_dto import SearchRequest, SearchResponse
from app.schemas.search_dto import AdvancedSearchFilter, SearchRequest, SearchResponse
from app.services.search_service import SearchService

router = APIRouter(prefix="/search", tags=["Search"])
Expand All @@ -17,12 +19,14 @@
async def search(request: SearchRequest, db: AsyncSession = Depends(get_db)) -> SearchResponse:
"""
Returns a list of papers that match the search query.

Currently, this returns the 5 most recently fetched papers, regardless
of the query string.
Optionally accepts an advanced filter with year range and text conditions.
"""

papers = await SearchService.search_papers(request.query, db)
papers = await SearchService.search_papers(
query=request.query,
db=db,
search_filter=request.filter,
)
return SearchResponse.model_validate(papers)


Expand All @@ -33,27 +37,47 @@ async def search(request: SearchRequest, db: AsyncSession = Depends(get_db)) ->
summary="Search for papers using a PDF",
)
async def search_by_pdf(
pdf: UploadFile = File(..., description="Research paper PDF"),
# Optional: user can also input a query
query: str | None = Form(
default=None,
description="Optional: query specifying what you want to find in relation to the paper",
),
db: AsyncSession = Depends(get_db),
pdf: UploadFile = File(..., description="Research paper PDF"),
query: str | None = Form(
default=None,
description="Optional: query specifying what you want to find in relation to the paper",
),
advanced_filter: str | None = Form(
default=None,
description="Optional: JSON-encoded advanced search filter",
),
db: AsyncSession = Depends(get_db),
) -> SearchResponse:
"""
Returns a list of papers that are relevant to the uploaded PDF
Returns a list of papers that are relevant to the uploaded PDF.
The PDF is analyzed, turned into semantic search queries, and used for vector search on our DB.
Optionally accepts a JSON-encoded advanced filter.
"""
if pdf.content_type != "application/pdf":
raise HTTPException(
status_code=400,
detail="Invalid file type: Only PDF files are supported.",
)

search_filter = None
if advanced_filter:
try:
search_filter = AdvancedSearchFilter.model_validate(json.loads(advanced_filter))
except json.JSONDecodeError as exc:
raise HTTPException(
status_code=400,
detail="Invalid filter JSON.",
) from exc
except ValueError as exc:
raise HTTPException(
status_code=400,
detail="Invalid filter.",
) from exc

papers = await SearchService.search_papers_from_pdf(
pdf_file=pdf,
db=db,
query=query,
search_filter=search_filter,
)
return SearchResponse.model_validate(papers)
42 changes: 40 additions & 2 deletions backend/app/schemas/search_dto.py
Original file line number Diff line number Diff line change
@@ -1,16 +1,54 @@
from typing import List
from __future__ import annotations

from pydantic import BaseModel
from typing import Annotated, List, Literal, Optional, Union

from pydantic import BaseModel, Field, model_validator

from app.schemas.paper_dto import PaperDto


class TextCondition(BaseModel):
"""Single text-based filter condition on a paper field."""

type: Literal["condition"]
field: Literal["title", "abstract"]
operator: Literal["contains", "not_contains"]
value: str


class ConditionGroup(BaseModel):
"""Logical group combining multiple conditions with AND / OR."""

type: Literal["group"]
operator: Literal["AND", "OR"]
children: List[
Annotated[Union[TextCondition, ConditionGroup], Field(discriminator="type")]
]


class AdvancedSearchFilter(BaseModel):
"""Structured filter with optional year range and boolean condition tree."""

year_from: Optional[int] = Field(default=None, ge=1, le=9999)
year_to: Optional[int] = Field(default=None, ge=1, le=9999)
root: ConditionGroup

@model_validator(mode="after")
def check_year_range(self) -> AdvancedSearchFilter:
"""Validate year_from is larger or equal to year_to."""
if self.year_from is not None and self.year_to is not None:
if self.year_from > self.year_to:
raise ValueError("year_from must be <= year_to")
return self


class SearchRequest(BaseModel):
"""
Request to search for specified query
"""

query: str
filter: Optional[AdvancedSearchFilter] = None


class SearchResponse(BaseModel):
Expand Down
28 changes: 19 additions & 9 deletions backend/app/services/search_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@

from app.core.deps import get_openai_provider, get_specter2_query_embedder
from app.repositories.search_repository import SearchRepository
from app.schemas.search_dto import PaperDto, SearchResponse
from app.schemas.search_dto import AdvancedSearchFilter, PaperDto, SearchResponse
from app.utils.author_utils import normalize_authors
from app.utils.pdf_utils import pdf_bytes_to_text
from app.utils.token_utils import ensure_fits_token_limit
Expand All @@ -24,7 +24,11 @@ class SearchService:
MAX_PDF_KEYWORD_INPUT_TOKENS = 280_000

@staticmethod
async def search_papers(query: str, db: AsyncSession) -> SearchResponse:
async def search_papers(
query: str,
db: AsyncSession,
search_filter: Optional[AdvancedSearchFilter] = None,
) -> SearchResponse:
"""
Search using a free-text query
"""
Expand All @@ -42,13 +46,15 @@ async def search_papers(query: str, db: AsyncSession) -> SearchResponse:
keywords=keywords,
db=db,
user_query=query,
search_filter=search_filter,
)

@staticmethod
async def search_papers_from_pdf(
pdf_file: UploadFile,
db: AsyncSession,
query: Optional[str] = None,
pdf_file: UploadFile,
db: AsyncSession,
query: Optional[str] = None,
search_filter: Optional[AdvancedSearchFilter] = None,
) -> SearchResponse:
"""
Search using a PDF as the primary signal.
Expand Down Expand Up @@ -97,15 +103,18 @@ async def search_papers_from_pdf(
logger.info("PDF search keywords: %s", keywords)

label = query or pdf_file.filename or "pdf-search"
return await SearchService._search_with_keywords(keywords=keywords, db=db, user_query=label)
return await SearchService._search_with_keywords(
keywords=keywords, db=db, user_query=label, search_filter=search_filter
)

# ---------- Shared search pipeline ----------

@staticmethod
async def _search_with_keywords(
keywords: List[str],
db: AsyncSession,
user_query: str,
keywords: List[str],
db: AsyncSession,
user_query: str,
search_filter: Optional[AdvancedSearchFilter] = None,
) -> SearchResponse:
"""
Core embedding + vector-search + DTO mapping pipeline.
Expand All @@ -122,6 +131,7 @@ async def _search_with_keywords(
db=db,
embeddings=embeddings,
limit=10,
search_filter=search_filter,
)

results: List[PaperDto] = []
Expand Down
6 changes: 5 additions & 1 deletion frontend/src/api/.openapi-generator/FILES
Original file line number Diff line number Diff line change
Expand Up @@ -8,10 +8,12 @@ base.ts
common.ts
configuration.ts
index.ts
models/advanced-search-filter.ts
models/chat-message-dto.ts
models/condition-group-children-inner.ts
models/condition-group.ts
models/httpvalidation-error.ts
models/index.ts
models/location-inner.ts
models/login-request.ts
Comment thread
MoSchmidt marked this conversation as resolved.
models/login-response.ts
models/paper-chat-request.ts
Expand All @@ -27,6 +29,8 @@ models/refresh-request.ts
models/refresh-response.ts
models/search-request.ts
models/search-response.ts
models/text-condition.ts
models/user-create.ts
models/user-response.ts
models/validation-error-loc-inner.ts
models/validation-error.ts
2 changes: 1 addition & 1 deletion frontend/src/api/.openapi-generator/VERSION
Original file line number Diff line number Diff line change
@@ -1 +1 @@
7.19.0
7.18.0-SNAPSHOT
2 changes: 1 addition & 1 deletion frontend/src/api/apis/authentication-api.ts
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ import type { AxiosPromise, AxiosInstance, RawAxiosRequestConfig } from 'axios';
import globalAxios from 'axios';
// Some imports not used depending on template conditions
// @ts-ignore
import { DUMMY_BASE_URL, assertParamExists, setApiKeyToObject, setBasicAuthToObject, setBearerAuthToObject, setOAuthToObject, setSearchParams, serializeDataIfNeeded, toPathString, createRequestFunction, replaceWithSerializableTypeIfNeeded } from '../common';
import { DUMMY_BASE_URL, assertParamExists, setApiKeyToObject, setBasicAuthToObject, setBearerAuthToObject, setOAuthToObject, setSearchParams, serializeDataIfNeeded, toPathString, createRequestFunction } from '../common';
Comment thread
MoSchmidt marked this conversation as resolved.
// @ts-ignore
import { BASE_PATH, COLLECTION_FORMATS, type RequestArgs, BaseAPI, RequiredError, operationServerMap } from '../base';
// @ts-ignore
Expand Down
2 changes: 1 addition & 1 deletion frontend/src/api/apis/paper-api.ts
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ import type { AxiosPromise, AxiosInstance, RawAxiosRequestConfig } from 'axios';
import globalAxios from 'axios';
// Some imports not used depending on template conditions
// @ts-ignore
import { DUMMY_BASE_URL, assertParamExists, setApiKeyToObject, setBasicAuthToObject, setBearerAuthToObject, setOAuthToObject, setSearchParams, serializeDataIfNeeded, toPathString, createRequestFunction, replaceWithSerializableTypeIfNeeded } from '../common';
import { DUMMY_BASE_URL, assertParamExists, setApiKeyToObject, setBasicAuthToObject, setBearerAuthToObject, setOAuthToObject, setSearchParams, serializeDataIfNeeded, toPathString, createRequestFunction } from '../common';
Comment thread
MoSchmidt marked this conversation as resolved.
// @ts-ignore
import { BASE_PATH, COLLECTION_FORMATS, type RequestArgs, BaseAPI, RequiredError, operationServerMap } from '../base';
// @ts-ignore
Expand Down
2 changes: 1 addition & 1 deletion frontend/src/api/apis/projects-api.ts
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@ import type { AxiosPromise, AxiosInstance, RawAxiosRequestConfig } from 'axios';
import globalAxios from 'axios';
// Some imports not used depending on template conditions
// @ts-ignore
import { DUMMY_BASE_URL, assertParamExists, setApiKeyToObject, setBasicAuthToObject, setBearerAuthToObject, setOAuthToObject, setSearchParams, serializeDataIfNeeded, toPathString, createRequestFunction, replaceWithSerializableTypeIfNeeded } from '../common';
import { DUMMY_BASE_URL, assertParamExists, setApiKeyToObject, setBasicAuthToObject, setBearerAuthToObject, setOAuthToObject, setSearchParams, serializeDataIfNeeded, toPathString, createRequestFunction } from '../common';
Comment thread
MoSchmidt marked this conversation as resolved.
// @ts-ignore
import { BASE_PATH, COLLECTION_FORMATS, type RequestArgs, BaseAPI, RequiredError, operationServerMap } from '../base';
// @ts-ignore
Expand Down
Loading