"""Deterministic query refinement from explicit relevance feedback."""
from __future__ import annotations
from collections import Counter
from collections.abc import Sequence
from typing import Protocol
from evidence_fetcher.models import Evidence, FeedbackLabel, RelevanceFeedback
from evidence_fetcher.ranking import (
_document_text,
_document_vector,
_inverse_document_frequency,
_tokenize,
)
MAX_EXPANSION_TERMS = 4
_STOP_WORDS = frozenset(
{
"a",
"an",
"and",
"are",
"as",
"at",
"be",
"by",
"for",
"from",
"how",
"in",
"is",
"it",
"of",
"on",
"or",
"that",
"the",
"this",
"to",
"with",
}
)
[docs]
class QueryRefiner(Protocol):
"""Produce a provider search query from evidence and relevance feedback."""
def refine(
self,
query: str,
evidence: Sequence[Evidence],
feedback: Sequence[RelevanceFeedback],
) -> str:
"""Return a refined provider query."""
[docs]
class TfidfQueryRefiner:
"""Small lexical query refiner based on feedback-marked evidence."""
[docs]
def __init__(self, *, max_terms: int = MAX_EXPANSION_TERMS) -> None:
if max_terms < 1:
raise ValueError("max_terms must be at least 1")
self._max_terms = max_terms
def refine(
self,
query: str,
evidence: Sequence[Evidence],
feedback: Sequence[RelevanceFeedback],
) -> str:
"""Append deterministic expansion terms to the original query."""
original_terms = _useful_terms(query)
material_feedback = [
item
for item in feedback
if item.label in {FeedbackLabel.MORE, FeedbackLabel.RELEVANT}
]
if not material_feedback:
raise ValueError("relevant feedback is needed to refine the query")
evidence_by_url = {item.url: item for item in evidence}
documents = [_tokenize(_document_text(item)) for item in evidence]
idf = _inverse_document_frequency(documents)
positive_weights: Counter[str] = Counter()
negative_weights: Counter[str] = Counter()
for item in feedback:
evidence_item = evidence_by_url.get(item.url)
if evidence_item is None:
raise ValueError("feedback URL is not present in the evidence set")
weights = _term_weights(evidence_item, idf)
match item.label:
case FeedbackLabel.MORE | FeedbackLabel.RELEVANT:
positive_weights.update(weights)
case FeedbackLabel.LESS:
negative_weights.update(
{term: weight * 0.25 for term, weight in weights.items()}
)
case FeedbackLabel.IRRELEVANT:
negative_weights.update(weights)
case FeedbackLabel.UNSURE:
continue
candidates = []
for term, weight in positive_weights.items():
if term in original_terms:
continue
negative_weight = negative_weights.get(term, 0.0)
if negative_weight >= weight:
continue
candidates.append((term, weight - negative_weight))
candidates.sort(key=lambda item: (-item[1], item[0]))
expansion_terms = [term for term, _ in candidates[: self._max_terms]]
if not expansion_terms:
raise ValueError("relevant feedback did not produce expansion terms")
return " ".join([query.strip(), *expansion_terms])
def _term_weights(evidence: Evidence, idf: dict[str, float]) -> dict[str, float]:
tokens = [
term
for term in _tokenize(_document_text(evidence))
if term not in _STOP_WORDS and len(term) > 2
]
return _document_vector(Counter(tokens), idf)
def _useful_terms(text: str) -> set[str]:
return {
term
for term in _tokenize(text)
if term not in _STOP_WORDS and len(term) > 2
}