Compare commits
No commits in common. "8d9b0d5f6315cff34785d7a6252e07162499ef94" and "f6e53758f7c089fa2f19f254ef25c8adc3571c00" have entirely different histories.
8d9b0d5f63
...
f6e53758f7
4 changed files with 45 additions and 110 deletions
|
|
@ -1 +0,0 @@
|
||||||
Subproject commit 3b9ff54ccf1540f8259a438b52d5bb05553e9dcb
|
|
||||||
|
|
@ -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_LOGIN = os.getenv("OPEN_API_LOGIN")
|
||||||
OPEN_API_PASSWORD = os.getenv("OPEN_API_PASSWORD")
|
OPEN_API_PASSWORD = os.getenv("OPEN_API_PASSWORD")
|
||||||
|
|
||||||
DENSE_PREFETCH_K = 80
|
DENSE_PREFETCH_K = 50
|
||||||
SPARSE_PREFETCH_K = 200
|
SPARSE_PREFETCH_K = 100
|
||||||
RETRIEVE_K = 150
|
RETRIEVE_K = 80
|
||||||
RERANK_LIMIT = 15
|
RERANK_LIMIT = 60
|
||||||
TOP_K = 50
|
TOP_K = 50
|
||||||
|
|
||||||
HTTP_TIMEOUT = 30.0
|
HTTP_TIMEOUT = 30.0
|
||||||
|
|
|
||||||
124
search/main.py
124
search/main.py
|
|
@ -162,11 +162,10 @@ async def lifespan(app: FastAPI):
|
||||||
|
|
||||||
app = FastAPI(title="Search Service", version="0.1.0", lifespan=lifespan)
|
app = FastAPI(title="Search Service", version="0.1.0", lifespan=lifespan)
|
||||||
|
|
||||||
DENSE_PREFETCH_K = 120
|
DENSE_PREFETCH_K = 80
|
||||||
SPARSE_PREFETCH_K = 200
|
SPARSE_PREFETCH_K = 200
|
||||||
RETRIEVE_K = 150
|
RETRIEVE_K = 150
|
||||||
RERANK_LIMIT = 35
|
RERANK_LIMIT = 15
|
||||||
KEYWORD_BOOST_EXTRA = 10
|
|
||||||
|
|
||||||
|
|
||||||
async def embed_dense(client: httpx.AsyncClient, text: str) -> list[float]:
|
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:
|
def build_dense_query(question: Question) -> str:
|
||||||
q = question.search_text.strip() if question.search_text else question.text.strip()
|
return question.text.strip()
|
||||||
return q
|
|
||||||
|
|
||||||
|
|
||||||
def build_sparse_query(question: Question) -> str:
|
def build_sparse_query(question: Question) -> str:
|
||||||
base = question.search_text.strip() if question.search_text else question.text.strip()
|
parts = [question.text.strip()]
|
||||||
parts = [base]
|
|
||||||
if question.keywords:
|
if question.keywords:
|
||||||
parts.extend(question.keywords)
|
parts.extend(question.keywords)
|
||||||
|
if question.search_text:
|
||||||
|
parts = [question.search_text]
|
||||||
return " ".join(parts)
|
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(
|
async def qdrant_search(
|
||||||
client: AsyncQdrantClient,
|
client: AsyncQdrantClient,
|
||||||
dense_vectors: list[list[float]],
|
dense_vectors: list[list[float]],
|
||||||
|
|
@ -362,24 +315,25 @@ async def rerank_points(
|
||||||
query: str,
|
query: str,
|
||||||
points: list[Any],
|
points: list[Any],
|
||||||
) -> list[Any]:
|
) -> list[Any]:
|
||||||
if not points:
|
rerank_candidates = points[:RERANK_LIMIT]
|
||||||
return []
|
rerank_targets = [point.payload.get("page_content") for point in rerank_candidates]
|
||||||
targets = [point.payload.get("page_content") for point in points]
|
scores = await get_rerank_scores(client, query, rerank_targets)
|
||||||
scores = await get_rerank_scores(client, query, targets)
|
|
||||||
|
|
||||||
if not scores or len(scores) != len(points):
|
if not scores:
|
||||||
logger.warning("Reranker unavailable or score mismatch, returning RRF order")
|
logger.warning("Reranker unavailable, returning RRF order")
|
||||||
return points
|
return rerank_candidates
|
||||||
|
|
||||||
return [
|
reranked_candidates = [
|
||||||
point
|
point
|
||||||
for _, point in sorted(
|
for _, point in sorted(
|
||||||
zip(scores, points),
|
zip(scores, rerank_candidates, strict=True),
|
||||||
key=lambda item: item[0],
|
key=lambda item: item[0],
|
||||||
reverse=True,
|
reverse=True,
|
||||||
)
|
)
|
||||||
]
|
]
|
||||||
|
|
||||||
|
return reranked_candidates
|
||||||
|
|
||||||
|
|
||||||
@app.get("/health")
|
@app.get("/health")
|
||||||
async def health() -> dict[str, str]:
|
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_vector, sparse_vector = await asyncio.gather(dense_task, sparse_task)
|
||||||
|
|
||||||
dense_vectors = [dense_vector]
|
dense_vectors = [dense_vector]
|
||||||
extra_texts: list[str] = []
|
if question.hyde and len(question.hyde) > 0:
|
||||||
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:
|
|
||||||
try:
|
try:
|
||||||
extra_vecs = await embed_dense_batch(client, extra_texts)
|
hyde_texts = question.hyde[:2]
|
||||||
dense_vectors.extend(extra_vecs)
|
hyde_vectors = await embed_dense_batch(client, hyde_texts)
|
||||||
|
dense_vectors.extend(hyde_vectors)
|
||||||
except Exception as e:
|
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)
|
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)
|
all_points = list(all_points)
|
||||||
|
|
||||||
rerank_pool, rerank_tail = prefilter_for_rerank(all_points, question)
|
reranked = await rerank_points(client, query, all_points)
|
||||||
reranked = await rerank_points(client, query, rerank_pool)
|
|
||||||
final_points = reranked + rerank_tail
|
|
||||||
|
|
||||||
msg_score: dict[str, float] = {}
|
reranked_ids = {id(p) for p in reranked}
|
||||||
for rank, point in enumerate(reranked):
|
remaining = [p for p in all_points if id(p) not in reranked_ids]
|
||||||
score = 1.0 / (rank + 1)
|
final_points = reranked + remaining
|
||||||
|
|
||||||
|
seen: set[str] = set()
|
||||||
|
message_ids: list[str] = []
|
||||||
|
for point in final_points:
|
||||||
for mid in extract_message_ids(point):
|
for mid in extract_message_ids(point):
|
||||||
msg_score[mid] = msg_score.get(mid, 0.0) + score
|
if mid not in seen:
|
||||||
for rank, point in enumerate(rerank_tail):
|
seen.add(mid)
|
||||||
score = 1.0 / (60 + rank + 1)
|
message_ids.append(mid)
|
||||||
for mid in extract_message_ids(point):
|
message_ids = message_ids[:50]
|
||||||
msg_score[mid] = msg_score.get(mid, 0.0) + score
|
|
||||||
message_ids = sorted(msg_score, key=lambda m: msg_score[m], reverse=True)[:50]
|
|
||||||
|
|
||||||
return SearchAPIResponse(results=[SearchAPIItem(message_ids=message_ids)])
|
return SearchAPIResponse(results=[SearchAPIItem(message_ids=message_ids)])
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -18,21 +18,15 @@ def _build_filter(question: Question) -> models.Filter | None:
|
||||||
must_conditions: list[models.Condition] = []
|
must_conditions: list[models.Condition] = []
|
||||||
|
|
||||||
if question.date_range:
|
if question.date_range:
|
||||||
try:
|
|
||||||
must_conditions.append(
|
|
||||||
models.FieldCondition(
|
|
||||||
key="metadata.end",
|
|
||||||
range=models.Range(gte=question.date_range.from_),
|
|
||||||
)
|
|
||||||
)
|
|
||||||
must_conditions.append(
|
must_conditions.append(
|
||||||
models.FieldCondition(
|
models.FieldCondition(
|
||||||
key="metadata.start",
|
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:
|
if question.asker:
|
||||||
must_conditions.append(
|
must_conditions.append(
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue