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}