import html import os import re import time from concurrent.futures import ThreadPoolExecutor from functools import lru_cache from dotenv import load_dotenv from openai import OpenAI from pinecone import Pinecone import fts_queries as fq load_dotenv() INDEX_NAME = "sec-fts" NAMESPACE = "__default__" EMBED_MODEL = "text-embedding-3-small" TICKERS = ["aapl", "amzn", "f", "gm", "msft", "orcl"] YEARS = list(range(2019, 2025)) INCLUDE_FIELDS = ["text", "ticker", "filing_type", "year", "chunk_index"] RRF_K = 60 def missing_env() -> list[str]: return [k for k in ("PINECONE_API_KEY", "OPENAI_API_KEY") if not os.environ.get(k)] @lru_cache(maxsize=1) def get_index(): return Pinecone(api_key=os.environ["PINECONE_API_KEY"]).index(name=INDEX_NAME) @lru_cache(maxsize=1) def get_openai() -> OpenAI: return OpenAI(api_key=os.environ["OPENAI_API_KEY"]) @lru_cache(maxsize=512) def _embed(text: str) -> tuple[float, ...]: resp = get_openai().embeddings.create(model=EMBED_MODEL, input=[text]) return tuple(resp.data[0].embedding) def embed(text: str) -> list[float]: return list(_embed(text)) def metadata_filter(tickers: list[str], years: list[int]) -> dict | None: filt = {} if tickers: filt["ticker"] = {"$in": list(tickers)} if years: filt["year"] = {"$in": [int(y) for y in years]} return filt or None def and_filters(*filters: dict | None) -> dict | None: parts = [f for f in filters if f] if not parts: return None return parts[0] if len(parts) == 1 else {"$and": parts} def search(req: dict): return get_index().documents.search(**req) def dense_request(query: str, top_k: int, filt: dict | None = None) -> dict: req = { "namespace": NAMESPACE, "top_k": top_k, "score_by": [{"type": "dense_vector", "field": "embedding", "values": embed(query)}], "include_fields": INCLUDE_FIELDS, } if filt: req["filter"] = filt return req def run_search_text(query: str, tickers: list, years: list, top_k: int): req = { "namespace": NAMESPACE, "top_k": top_k, "score_by": [{"type": "text", "field": "text", "query": query}], "include_fields": INCLUDE_FIELDS, } filt = metadata_filter(tickers, years) if filt: req["filter"] = filt return search(req) def run_search_semantic(query: str, tickers: list, years: list, top_k: int): return search(dense_request(query, top_k, metadata_filter(tickers, years))) def run_search_hybrid(semantic_query: str, text_filter: str, tickers: list, years: list, top_k: int): filt = metadata_filter(tickers, years) or {} if text_filter.strip(): filt["text"] = {"$match_all": text_filter.strip()} return search(dense_request(semantic_query, top_k, filt or None)) def term_pattern(terms: list[str]) -> re.Pattern | None: if not terms: return None stems = [re.escape(t[:-2] if len(t) > 5 else t) for t in terms] return re.compile(r"\b(" + "|".join(stems) + r")\w*", re.IGNORECASE) def highlight(text: str, terms: list[str]) -> tuple[str, str]: escaped = html.escape(text) pattern = term_pattern(terms) if not pattern: return escaped[:300] + ("..." if len(escaped) > 300 else ""), escaped marked = pattern.sub(lambda m: f"{m.group(0)}", escaped) first = pattern.search(escaped) start = max(0, first.start() - 150) if first else 0 snippet_raw = escaped[start : start + 450] snippet = pattern.sub(lambda m: f"{m.group(0)}", snippet_raw) snippet = ("..." if start else "") + snippet + ("..." if start + 450 < len(escaped) else "") return snippet, marked def keyword_coverage(text: str, terms: list[str]) -> tuple[int, int]: hits = sum(1 for t in terms if term_pattern([t]).search(text)) return hits, len(terms) def match_record(m) -> dict: d = m.to_dict() return { "id": getattr(m, "_id", d.get("_id")), "score": getattr(m, "_score", getattr(m, "score", d.get("_score", 0.0))), "ticker": str(d.get("ticker", "?")).upper(), "year": int(d.get("year", 0)), "filing": str(d.get("filing_type", "")).upper(), "chunk": int(d["chunk_index"]) if d.get("chunk_index") is not None else "?", "text": str(d.get("text", "")), } def _timed(fn): start = time.perf_counter() result = fn() return result, (time.perf_counter() - start) * 1000 def run_compare(query: str, fts_state: dict, base_filter: dict | None, dense_extra_filter: dict | None, top_k: int) -> dict: fts_req = fq.build_request(fts_state, top_k, INCLUDE_FIELDS) fts_req["filter"] = and_filters(fts_req.get("filter"), base_filter) if not fts_req["filter"]: del fts_req["filter"] def dense(): _, embed_ms = _timed(lambda: embed(query)) resp, search_ms = _timed(lambda: search(dense_request(query, top_k, and_filters(base_filter, dense_extra_filter)))) return resp, embed_ms, search_ms with ThreadPoolExecutor(max_workers=2) as pool: dense_future = pool.submit(dense) fts_future = pool.submit(_timed, lambda: search(fts_req)) dense_resp, embed_ms, dense_ms = dense_future.result() fts_resp, fts_ms = fts_future.result() return { "dense": [match_record(m) for m in dense_resp.matches], "fts": [match_record(m) for m in fts_resp.matches], "embed_ms": embed_ms, "dense_ms": dense_ms, "fts_ms": fts_ms, "fts_req": fts_req, "terms": fq.highlight_terms(fts_state), } def compare_stats(res: dict) -> dict: dense, fts, terms = res["dense"], res["fts"], res["terms"] dense_rank = {r["id"]: i for i, r in enumerate(dense, 1)} fts_rank = {r["id"]: i for i, r in enumerate(fts, 1)} shared = dense_rank.keys() & fts_rank.keys() union = dense_rank.keys() | fts_rank.keys() def avg_coverage(rows): if not rows or not terms: return None return sum(keyword_coverage(r["text"], terms)[0] / len(terms) for r in rows) / len(rows) return { "dense_rank": dense_rank, "fts_rank": fts_rank, "shared": shared, "union": union, "by_id": {r["id"]: r for r in dense + fts}, "cov_dense": avg_coverage(dense), "cov_fts": avg_coverage(fts), "k": max(len(dense), len(fts)), } def rank_rows(res: dict, stats: dict) -> list[dict]: terms = res["terms"] rows = [] for doc_id in stats["union"]: r = stats["by_id"][doc_id] d, f = stats["dense_rank"].get(doc_id), stats["fts_rank"].get(doc_id) hits, n = keyword_coverage(r["text"], terms) rows.append({ "ticker": r["ticker"], "year": r["year"], "chunk": r["chunk"], "dense rank": d, "full-text rank": f, "Δ rank (dense − FTS)": d - f if d and f else None, "found by": "both" if d and f else ("dense" if d else "full-text"), "keywords": f"{hits}/{n}" if n else "", "snippet": r["text"][:140].replace("\n", " "), }) rows.sort(key=lambda x: min(x["dense rank"] or 999, x["full-text rank"] or 999)) return rows def rrf_fuse(stats: dict) -> list[tuple[str, float]]: fused = {} for ranks in (stats["dense_rank"], stats["fts_rank"]): for doc_id, rank in ranks.items(): fused[doc_id] = fused.get(doc_id, 0.0) + 1 / (RRF_K + rank) return sorted(fused.items(), key=lambda kv: kv[1], reverse=True)[: stats["k"]] def rank_badge(doc_id: str, other_ranks: dict, other_name: str) -> str: if doc_id in other_ranks: return f"🔁 also #{other_ranks[doc_id]} in {other_name}" return f"◆ only in {'dense' if other_name == 'full-text' else 'full-text'}"