Chat-Service / app /tools /hybrid_tool.py
ArabicNewsAnalyzer's picture
Upload 56 files
de28957 verified
Raw
History Blame Contribute Delete
5.29 kB
from __future__ import annotations
import json
from typing import Any
from langchain_core.tools import tool
from app.db.neo4j import run_read_query
from app.db.postgres import get_readonly_connection
from app.services.nlp_client import embed
# Must match FAISS_DIM / MiniLM output dim used by nlp_client.embed()
EMBEDDING_DIM = 384
def _to_vector_literal(vec) -> str:
"""
Format an embedding as a pgvector literal, e.g. '[0.1,0.2,...]'.
No pgvector adapter is registered on the psycopg pool (see
app/db/postgres.py), so we build the literal ourselves and cast
with ::vector in SQL rather than relying on parameter binding.
"""
return "[" + ",".join(f"{float(x):.8f}" for x in vec) + "]"
async def _semantic_search(query_text: str, limit: int) -> list[dict[str, Any]]:
"""Step 1 + 2: embed query_text, then pgvector cosine similarity search."""
vectors = await embed([query_text])
vector_literal = _to_vector_literal(vectors[0])
sql = """
SELECT id, title, url, published_at, content,
(search_vector <=> %s::vector) AS distance
FROM news_articles
WHERE search_vector IS NOT NULL
ORDER BY distance ASC
LIMIT %s
"""
async with get_readonly_connection() as conn:
cur = await conn.execute(sql, (vector_literal, limit))
rows = await cur.fetchall()
columns = [desc[0] for desc in cur.description]
articles = []
for row in rows:
record = dict(zip(columns, row))
content = record.pop("content", None) or ""
snippet = content[:280] + ("..." if len(content) > 280 else "")
distance = record.pop("distance")
articles.append({
"id": record["id"],
"title": record["title"],
"url": record["url"],
"published_at": (
record["published_at"].isoformat()
if record.get("published_at") else None
),
"snippet": snippet,
# cosine distance -> similarity; pgvector's <=> is cosine
# distance under the vector_cosine_ops opclass
"similarity": round(1 - distance, 4),
})
return articles
async def _entities_for_article(article_id: int, limit: int) -> list[dict[str, Any]]:
"""Step 3: pull mentioned entities from Neo4j for one article."""
query = """
MATCH (a:Article {article_id: $article_id})-[:MENTIONS]->(e:Entity)
RETURN e.name AS name, e.type AS type, e.entity_id AS entity_id
LIMIT $limit
"""
rows = await run_read_query(query, article_id=article_id, limit=limit)
return [r for r in rows if r.get("entity_id")]
@tool
async def hybrid_search_tool(query_text: str, limit: int = 10) -> str:
"""
Semantic search over article content, expanded with mentioned entities.
Use this when a question needs *finding* articles by meaning/topic
(not an exact stat or a known entity/article_id lookup) — e.g. "articles
about the election and who's mentioned in them", "coverage of the
drought", "find articles similar to X".
Steps:
1. Embed query_text (shared MiniLM endpoint, same as Django uses).
2. Cosine similarity search over news_articles.search_vector
(pgvector) to get the top-`limit` articles, ranked by similarity.
3. For each returned article, pull its mentioned entities from
Neo4j (:MENTIONS relationship).
4. Merge into one ranked result set with combined source_refs
(articles + entities).
query_text: free-text description of what to search for.
limit: max number of articles to return (entities per article are
capped separately and not counted against this).
Returns a JSON string: {"rows": [{"id", "title", "url", "published_at",
"snippet", "similarity", "entities": [...]}], "source_refs": [{"type":
"article"|"entity", "id": ...}, ...]}
"""
try:
articles = await _semantic_search(query_text, limit)
except Exception as exc: # noqa: BLE001 - surface to agent, don't crash turn
return json.dumps({"rows": [], "source_refs": [], "error": str(exc)})
source_refs: list[dict[str, Any]] = []
seen_entity_ids: set[Any] = set()
for article in articles:
source_refs.append({"type": "article", "id": article["id"]})
try:
entities = await _entities_for_article(article["id"], limit=10)
except Exception: # noqa: BLE001 - one article's graph lookup failing shouldn't sink the rest
entities = []
article["entities"] = [
{"name": e.get("name"), "type": e.get("type"),
"entity_id": e.get("entity_id")}
for e in entities
]
for e in entities:
entity_id = e.get("entity_id")
if entity_id is not None and entity_id not in seen_entity_ids:
seen_entity_ids.add(entity_id)
source_refs.append({"type": "entity", "id": entity_id})
# already ranked by similarity from the SQL ORDER BY
return json.dumps({"rows": articles, "source_refs": source_refs}, default=str)