Compare commits

..

No commits in common. "8d9b0d5f6315cff34785d7a6252e07162499ef94" and "f6e53758f7c089fa2f19f254ef25c8adc3571c00" have entirely different histories.

4 changed files with 45 additions and 110 deletions

@ -1 +0,0 @@
Subproject commit 3b9ff54ccf1540f8259a438b52d5bb05553e9dcb

View file

@ -19,10 +19,10 @@ QDRANT_SPARSE_VECTOR_NAME = os.getenv("QDRANT_SPARSE_VECTOR_NAME", "sparse")
OPEN_API_LOGIN = os.getenv("OPEN_API_LOGIN")
OPEN_API_PASSWORD = os.getenv("OPEN_API_PASSWORD")
DENSE_PREFETCH_K = 80
SPARSE_PREFETCH_K = 200
RETRIEVE_K = 150
RERANK_LIMIT = 15
DENSE_PREFETCH_K = 50
SPARSE_PREFETCH_K = 100
RETRIEVE_K = 80
RERANK_LIMIT = 60
TOP_K = 50
HTTP_TIMEOUT = 30.0

View file

@ -162,11 +162,10 @@ async def lifespan(app: FastAPI):
app = FastAPI(title="Search Service", version="0.1.0", lifespan=lifespan)
DENSE_PREFETCH_K = 120
DENSE_PREFETCH_K = 80
SPARSE_PREFETCH_K = 200
RETRIEVE_K = 150
RERANK_LIMIT = 35
KEYWORD_BOOST_EXTRA = 10
RERANK_LIMIT = 15
async def embed_dense(client: httpx.AsyncClient, text: str) -> list[float]:
@ -212,64 +211,18 @@ def embed_sparse_sync(text: str) -> SparseVector:
def build_dense_query(question: Question) -> str:
q = question.search_text.strip() if question.search_text else question.text.strip()
return q
return question.text.strip()
def build_sparse_query(question: Question) -> str:
base = question.search_text.strip() if question.search_text else question.text.strip()
parts = [base]
parts = [question.text.strip()]
if question.keywords:
parts.extend(question.keywords)
if question.search_text:
parts = [question.search_text]
return " ".join(parts)
def _build_keyword_set(question: Question) -> list[str]:
tokens: list[str] = []
if question.keywords:
tokens.extend(kw.lower() for kw in question.keywords if kw)
if question.entities:
for field in (
question.entities.people,
question.entities.emails,
question.entities.documents,
question.entities.names,
question.entities.links,
):
tokens.extend(e.lower() for e in (field or []) if e)
return tokens
def prefilter_for_rerank(
points: list[Any],
question: Question,
) -> tuple[list[Any], list[Any]]:
"""Select candidates for reranking: top by RRF + keyword-boosted stragglers."""
if not points:
return [], []
head = points[:RERANK_LIMIT]
tail = points[RERANK_LIMIT:]
keywords = _build_keyword_set(question)
if not keywords or not tail:
return head, tail
extra: list[Any] = []
remaining_tail: list[Any] = []
for p in tail:
if len(extra) >= KEYWORD_BOOST_EXTRA:
remaining_tail.append(p)
continue
content = ((p.payload or {}).get("page_content") or "").lower()
if any(kw in content for kw in keywords):
extra.append(p)
else:
remaining_tail.append(p)
return head + extra, remaining_tail
async def qdrant_search(
client: AsyncQdrantClient,
dense_vectors: list[list[float]],
@ -362,24 +315,25 @@ async def rerank_points(
query: str,
points: list[Any],
) -> list[Any]:
if not points:
return []
targets = [point.payload.get("page_content") for point in points]
scores = await get_rerank_scores(client, query, targets)
rerank_candidates = points[:RERANK_LIMIT]
rerank_targets = [point.payload.get("page_content") for point in rerank_candidates]
scores = await get_rerank_scores(client, query, rerank_targets)
if not scores or len(scores) != len(points):
logger.warning("Reranker unavailable or score mismatch, returning RRF order")
return points
if not scores:
logger.warning("Reranker unavailable, returning RRF order")
return rerank_candidates
return [
reranked_candidates = [
point
for _, point in sorted(
zip(scores, points),
zip(scores, rerank_candidates, strict=True),
key=lambda item: item[0],
reverse=True,
)
]
return reranked_candidates
@app.get("/health")
async def health() -> dict[str, str]:
@ -404,25 +358,13 @@ async def search(payload: SearchAPIRequest) -> SearchAPIResponse:
dense_vector, sparse_vector = await asyncio.gather(dense_task, sparse_task)
dense_vectors = [dense_vector]
extra_texts: list[str] = []
raw_text = question.text.strip()
if raw_text and raw_text != dense_query:
extra_texts.append(raw_text)
for v in (question.variants or []):
q_v = v.strip()
if q_v and q_v != dense_query and q_v not in extra_texts:
extra_texts.append(q_v)
for h in (question.hyde or []):
q_h = h.strip()
if q_h and q_h != dense_query and q_h not in extra_texts:
extra_texts.append(q_h)
extra_texts = extra_texts[:3]
if extra_texts:
if question.hyde and len(question.hyde) > 0:
try:
extra_vecs = await embed_dense_batch(client, extra_texts)
dense_vectors.extend(extra_vecs)
hyde_texts = question.hyde[:2]
hyde_vectors = await embed_dense_batch(client, hyde_texts)
dense_vectors.extend(hyde_vectors)
except Exception as e:
logger.warning(f"Extra dense embedding failed: {e}")
logger.warning(f"HyDE embedding failed: {e}")
all_points = await qdrant_search(qdrant, dense_vectors, sparse_vector)
@ -431,20 +373,20 @@ async def search(payload: SearchAPIRequest) -> SearchAPIResponse:
all_points = list(all_points)
rerank_pool, rerank_tail = prefilter_for_rerank(all_points, question)
reranked = await rerank_points(client, query, rerank_pool)
final_points = reranked + rerank_tail
reranked = await rerank_points(client, query, all_points)
msg_score: dict[str, float] = {}
for rank, point in enumerate(reranked):
score = 1.0 / (rank + 1)
reranked_ids = {id(p) for p in reranked}
remaining = [p for p in all_points if id(p) not in reranked_ids]
final_points = reranked + remaining
seen: set[str] = set()
message_ids: list[str] = []
for point in final_points:
for mid in extract_message_ids(point):
msg_score[mid] = msg_score.get(mid, 0.0) + score
for rank, point in enumerate(rerank_tail):
score = 1.0 / (60 + rank + 1)
for mid in extract_message_ids(point):
msg_score[mid] = msg_score.get(mid, 0.0) + score
message_ids = sorted(msg_score, key=lambda m: msg_score[m], reverse=True)[:50]
if mid not in seen:
seen.add(mid)
message_ids.append(mid)
message_ids = message_ids[:50]
return SearchAPIResponse(results=[SearchAPIItem(message_ids=message_ids)])

View file

@ -18,21 +18,15 @@ def _build_filter(question: Question) -> models.Filter | None:
must_conditions: list[models.Condition] = []
if question.date_range:
try:
must_conditions.append(
models.FieldCondition(
key="metadata.end",
range=models.Range(gte=question.date_range.from_),
)
)
must_conditions.append(
models.FieldCondition(
key="metadata.start",
range=models.Range(lte=question.date_range.to),
datetime_range=models.DatetimeRange(
gte=question.date_range.from_,
lte=question.date_range.to,
),
)
)
except Exception as e:
logger.warning("Date filter failed: %s", e)
if question.asker:
must_conditions.append(