Spaces:
Sleeping
Sleeping
File size: 7,468 Bytes
6ba3ef3 2e01d2b 6ba3ef3 2e01d2b 6ba3ef3 2e01d2b 6ba3ef3 8db8a8c 6ba3ef3 8db8a8c 48c9780 8db8a8c 6ba3ef3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 | """FinChat retrieval-augmented generation (RAG) chain.
Public entry point: answer(question) -> {"answer", "routed_to", "sources"}
"""
from __future__ import annotations
import os
import re
from collections import defaultdict
from functools import lru_cache
from dotenv import load_dotenv
from langchain_chroma import Chroma
from langchain_huggingface import HuggingFaceEmbeddings
from langchain_groq import ChatGroq
from langchain_core.prompts import ChatPromptTemplate
from src import config
# Load GROQ_API_KEY from the project's .env by explicit path (robust no matter
# where the process is launched). On Hugging Face Spaces there is no .env and
# this is a harmless no-op -- the key comes from Space secrets instead.
load_dotenv(config.PROJECT_ROOT / ".env")
SYSTEM_PROMPT = """You are FinChat, a financial analyst assistant. Answer the \
question using ONLY the excerpts from SEC 10-K filings provided below.
Rules:
- Use ONLY the provided context. Do NOT rely on outside knowledge.
- If the answer is not in the context, reply exactly:
"I couldn't find that in the filings I have."
- Be concise and precise with numbers. State the company and fiscal year.
- End with a short "Sources:" list referencing the excerpts you used.
Context:
{context}
"""
PROMPT = ChatPromptTemplate.from_messages(
[("system", SYSTEM_PROMPT), ("human", "{question}")]
)
@lru_cache(maxsize=1)
def get_vectorstore() -> Chroma:
embeddings = HuggingFaceEmbeddings(model_name=config.EMBEDDING_MODEL)
return Chroma(
collection_name=config.CHROMA_COLLECTION,
embedding_function=embeddings,
persist_directory=str(config.VECTORSTORE_DIR),
)
def ensure_index() -> None:
"""Build the vector store on first run if it doesn't exist yet.
Lets the app bootstrap itself on a fresh deployment (e.g. Hugging Face
Spaces). On normal runs where the store already exists, this is a fast
no-op.
IMPORTANT: this checks the filesystem instead of opening a Chroma client.
Probing with a client would create an empty database and keep it open --
and because chromadb caches clients per path, the subsequent rebuild
(delete + recreate the directory) would leave that cached client pointing
at deleted files, crashing with "unable to open database file".
"""
if not (config.VECTORSTORE_DIR / "chroma.sqlite3").exists():
from src.ingest import build_index
build_index()
@lru_cache(maxsize=1)
def get_llm() -> ChatGroq:
if not os.getenv("GROQ_API_KEY"):
raise RuntimeError("GROQ_API_KEY is not set. Add it to your .env file.")
return ChatGroq(
model=config.LLM_MODEL,
temperature=config.LLM_TEMPERATURE,
max_retries=5, # back off through Groq free-tier rate limits
)
# Words that shouldn't count as a company alias on their own.
_STOPWORDS = {
"inc", "corp", "corporation", "company", "ltd", "llc", "plc", "the",
"and", "group", "holdings", "international", "industries", "products",
"resources", "energy", "technologies", "systems",
}
@lru_cache(maxsize=1)
def company_aliases() -> dict[str, str]:
"""Map each recognizable alias (ticker or distinctive name word) -> ticker.
This lets FinChat route a question to the right company before retrieving
("knows exactly where to look"). Any alias shared by more than one company
is dropped, so we never route to the wrong filing.
"""
store = get_vectorstore()
metadatas = store.get(include=["metadatas"]).get("metadatas", [])
ticker_to_name: dict[str, str] = {}
for m in metadatas:
ticker = (m.get("ticker") or "").upper()
if ticker:
ticker_to_name[ticker] = m.get("company") or ""
alias_to_tickers: dict[str, set] = defaultdict(set)
for ticker, name in ticker_to_name.items():
alias_to_tickers[ticker.lower()].add(ticker) # the ticker itself
for word in re.findall(r"[a-z]+", name.lower()):
if len(word) >= 4 and word not in _STOPWORDS:
alias_to_tickers[word].add(ticker)
# Keep only unambiguous aliases (mapping to exactly one company).
return {a: next(iter(ts)) for a, ts in alias_to_tickers.items() if len(ts) == 1}
@lru_cache(maxsize=1)
def available_companies() -> list[tuple[str, str]]:
"""Return sorted (ticker, company name) pairs present in the vector store."""
store = get_vectorstore()
metadatas = store.get(include=["metadatas"]).get("metadatas", [])
seen: dict[str, str] = {}
for m in metadatas:
ticker = (m.get("ticker") or "").upper()
if ticker and ticker not in seen:
seen[ticker] = m.get("company") or ticker
return sorted(seen.items())
def detect_ticker(question: str) -> str | None:
"""Figure out which company the question is about."""
aliases = company_aliases()
# 1) An explicit ticker written in capitals, e.g. "AMD" or "ABT".
for token in re.findall(r"\b[A-Z]{2,6}\b", question):
if token.lower() in aliases:
return aliases[token.lower()]
# 2) A distinctive company-name word, e.g. "abbott" or "matson".
for word in re.findall(r"[a-z]+", question.lower()):
if len(word) >= 4 and word in aliases:
return aliases[word]
return None
# Terms that signal a numeric/financial question -> pull in the XBRL statements.
_FINANCIAL_TERMS = (
"revenue", "sales", "income", "earnings", "profit", "margin", "ebitda",
"asset", "liabilit", "equity", "cash flow", "cash", "debt", "expense",
"eps", "per share", "how much", "dividend", "operating", "gross", "net ",
"balance sheet", "capital", "ratio",
)
def _is_financial_query(question: str) -> bool:
q = question.lower()
return any(term in q for term in _FINANCIAL_TERMS)
def retrieve(question: str, ticker: str | None):
store = get_vectorstore()
search_kwargs: dict = {"k": config.TOP_K}
if ticker:
# Metadata filter = search ONLY that company's filings.
search_kwargs["filter"] = {"ticker": ticker}
docs = store.as_retriever(search_kwargs=search_kwargs).invoke(question)
# Hybrid step: for a numeric/financial question about a known company,
# guarantee that company's structured XBRL statements are in context --
# they can otherwise be out-ranked by revenue *discussion* in the filing
# text (as happens for Apple).
if ticker and _is_financial_query(question):
fin = store.as_retriever(
search_kwargs={
"k": 4,
"filter": {"$and": [{"ticker": ticker}, {"type": "financials"}]},
}
).invoke(question)
seen = {d.page_content[:80] for d in fin}
rest = [d for d in docs if d.page_content[:80] not in seen]
docs = (fin + rest)[: config.TOP_K]
return docs
def format_context(docs) -> str:
return "\n\n".join(
f"[{i}] {d.metadata.get('source', 'source')}\n{d.page_content}"
for i, d in enumerate(docs, 1)
)
def answer(question: str) -> dict:
"""Route -> retrieve -> generate. Returns answer, routing info, sources."""
ticker = detect_ticker(question)
docs = retrieve(question, ticker)
context = format_context(docs) if docs else "(no relevant excerpts found)"
chain = PROMPT | get_llm()
response = chain.invoke({"context": context, "question": question})
return {"answer": response.content, "routed_to": ticker, "sources": docs}
|