study-buddy / app /services /literature_contract.py
GitHub Actions
deploy d092bea3608b7a29952f16357fda39b7a29e399b
2e818da
Raw
History Blame Contribute Delete
5.06 kB
from __future__ import annotations
import re
from typing import Literal
from pydantic import BaseModel, ConfigDict, Field
LiteratureOperation = Literal[
"exact_paper",
"author_works",
"recommend",
"citing_papers",
"referenced_papers",
"ranked_papers",
"literature_review",
]
LiteratureRanking = Literal[
"relevance",
"citation_count",
"publication_date",
"within_year_citation",
"combined",
]
CitationDirection = Literal["none", "cites_seed", "cited_by_seed"]
SelectionMode = Literal["top_overall", "one_per_year"]
class LiteratureRequest(BaseModel):
model_config = ConfigDict(extra="forbid")
operation: LiteratureOperation = "recommend"
paper_anchors: list[str] = Field(default_factory=list)
author_names: list[str] = Field(default_factory=list)
topics: list[str] = Field(default_factory=list)
search_queries: list[str] = Field(default_factory=list)
citation_direction: CitationDirection = "none"
ranking: LiteratureRanking = "combined"
year_start: int | None = None
year_end: int | None = None
selection_mode: SelectionMode = "top_overall"
result_limit: int = Field(default=5, ge=1, le=50)
clarification_attempted: bool = False
assumptions: list[str] = Field(default_factory=list)
_COUNT_WORDS = {
"one": 1, "two": 2, "three": 3, "four": 4, "five": 5,
"six": 6, "seven": 7, "eight": 8, "nine": 9, "ten": 10,
"eleven": 11, "twelve": 12, "fifteen": 15, "twenty": 20,
"thirty": 30, "forty": 40, "fifty": 50,
}
def _clean(value: str) -> str:
return " ".join(value.strip(" .,!?;:\"'“”").split())
def _explicit_count(query: str) -> int | None:
match = re.search(
r"\b(?:top|best|find|show|give|recommend)?\s*"
r"(\d{1,2}|one|two|three|four|five|six|seven|eight|nine|ten|eleven|twelve|fifteen|twenty|thirty|forty|fifty)"
r"\s+(?:most\s+)?(?:cited\s+)?papers?\b",
query,
re.I,
)
if not match:
return None
raw = match.group(1).casefold()
value = int(raw) if raw.isdigit() else _COUNT_WORDS.get(raw)
return max(1, min(50, value)) if value is not None else None
def _author_name(query: str) -> str | None:
match = re.search(
r"\b(?:latest|newest|recent|top|best)?\s*papers?\s+(?:by|of|from)\s+"
r"([A-Z][\w.'’-]+(?:\s+[A-Z][\w.'’-]+){1,4})(?=\s*(?:[?.!,;]|$|\b(?:in|from|during|published)\s+(?:19|20)\d{2}\b))",
query,
)
return _clean(match.group(1)) if match else None
def _citation_shape(query: str) -> tuple[CitationDirection, str | None]:
normalized = " ".join(query.split())
incoming = re.search(r"\bpapers?\s+that\s+(?:cite|cited|reference|referenced)\s+(.+?)(?:[?.!,;]|$)", normalized, re.I)
if incoming:
return "cites_seed", _clean(incoming.group(1))
outgoing = re.search(r"\bpapers?\s+(.+?)\s+(?:cites|references)\s*(?:[?.!,;]|$)", normalized, re.I)
if outgoing:
return "cited_by_seed", _clean(outgoing.group(1))
return "none", None
def build_deterministic_contract(
query: str,
planned: LiteratureRequest,
*,
selection_text: str | None = None,
) -> LiteratureRequest:
result = planned.model_copy(deep=True)
text = " ".join(query.split()).strip()
lowered = text.casefold()
explicit_count = _explicit_count(text)
author = _author_name(text)
direction, citation_anchor = _citation_shape(text)
years = [int(value) for value in re.findall(r"\b(?:19\d{2}|20\d{2})\b", text)]
if author:
result.operation = "author_works"
result.author_names = [author]
if direction != "none":
result.citation_direction = direction
result.operation = "citing_papers" if direction == "cites_seed" else "referenced_papers"
if citation_anchor:
result.paper_anchors = [citation_anchor]
if re.search(r"\b(?:literature review|review the literature|survey papers?|state of the art)\b", lowered):
result.operation = "literature_review"
elif not author and direction == "none" and re.search(r"\btop\b.*\bpapers?\b|\bpapers?.*\bby citations?\b", lowered):
result.operation = "ranked_papers"
if re.search(r"\b(?:latest|newest|most recent)\b", lowered):
result.ranking = "publication_date"
elif re.search(r"\b(?:by citations?|most cited|citation count)\b", lowered) or direction != "none" and re.search(r"\btop\b", lowered):
result.ranking = "citation_count"
if years:
result.year_start = min(years)
result.year_end = max(years)
if explicit_count is not None:
result.result_limit = explicit_count
elif author and re.search(r"\b(?:the\s+)?(?:latest|newest|most recent)\s+paper\b", lowered):
result.result_limit = 1
elif result.operation == "exact_paper":
result.result_limit = 1
if not result.paper_anchors and selection_text and direction != "none":
selected = _clean(selection_text[:300])
if selected:
result.paper_anchors = [selected]
return result